From eba5496019c6a1a5d84e7bd39f69a9bcef502d47 Mon Sep 17 00:00:00 2001 From: Kt Date: Thu, 24 Sep 2026 14:35:17 +0800 Subject: [PATCH] refactor(session)!: extract in-memory backend package Move the memory service, views, and tests to session/inmemory. Preserve storage behavior and launcher defaults. BREAKING CHANGE: replace session.InMemoryService() with inmemory.New(). --- agent/llmagent/agent_test.go | 3 +- cmd/launcher/config.go | 3 +- cmd/launcher/console/console_test.go | 3 +- cmd/launcher/web/web_test.go | 3 +- runner/runner_test.go | 3 +- server/server_test.go | 7 +- session/inmemory.go | 193 ------------------ session/inmemory/service.go | 131 ++++++++++++ .../service_test.go} | 94 ++++++++- session/inmemory/session.go | 74 +++++++ session/session_test.go | 59 +----- 11 files changed, 311 insertions(+), 262 deletions(-) delete mode 100644 session/inmemory.go create mode 100644 session/inmemory/service.go rename session/{inmemory_test.go => inmemory/service_test.go} (76%) create mode 100644 session/inmemory/session.go diff --git a/agent/llmagent/agent_test.go b/agent/llmagent/agent_test.go index 6c6ba54..2260d70 100644 --- a/agent/llmagent/agent_test.go +++ b/agent/llmagent/agent_test.go @@ -16,6 +16,7 @@ import ( "github.com/sylumi/agentkit/agent/llmagent" "github.com/sylumi/agentkit/model" "github.com/sylumi/agentkit/session" + "github.com/sylumi/agentkit/session/inmemory" "github.com/sylumi/agentkit/tool" "github.com/sylumi/agentkit/tool/functiontool" ) @@ -152,7 +153,7 @@ func toolCallResult(calls ...model.ToolCallPart) model.Result { func newInvocation(t *testing.T) (session.Service, *agent.InvocationContext) { t.Helper() - service := session.InMemoryService() + service := inmemory.New() created, err := service.Create(t.Context(), &session.CreateRequest{AppName: "test", UserID: "user"}) if err != nil { t.Fatal(err) diff --git a/cmd/launcher/config.go b/cmd/launcher/config.go index 858191b..0269f89 100644 --- a/cmd/launcher/config.go +++ b/cmd/launcher/config.go @@ -7,6 +7,7 @@ import ( "github.com/sylumi/agentkit/agent" "github.com/sylumi/agentkit/session" + "github.com/sylumi/agentkit/session/inmemory" ) // Config supplies an Agent and optional services to console and web launchers. @@ -35,7 +36,7 @@ func (c Config) Resolve() (Config, error) { return Config{}, fmt.Errorf("launcher: app name and user ID must not be blank") } if c.SessionService == nil { - c.SessionService = session.InMemoryService() + c.SessionService = inmemory.New() } return c, nil } diff --git a/cmd/launcher/console/console_test.go b/cmd/launcher/console/console_test.go index 641db2d..56c75b0 100644 --- a/cmd/launcher/console/console_test.go +++ b/cmd/launcher/console/console_test.go @@ -14,6 +14,7 @@ import ( "github.com/sylumi/agentkit/cmd/launcher" "github.com/sylumi/agentkit/model" "github.com/sylumi/agentkit/session" + "github.com/sylumi/agentkit/session/inmemory" ) type testAgent struct { @@ -36,7 +37,7 @@ func TestConsoleKeepsHistoryAndCancelsWhileWaitingForInput(t *testing.T) { yield(e, nil) } }} - store := session.InMemoryService() + store := inmemory.New() cfg := launcher.Config{Agent: a, SessionService: store, AppName: "custom-app", UserID: "alice"} var output bytes.Buffer if err := run(t.Context(), cfg, strings.NewReader("hi\n\nagain\n"), &output); err != nil { diff --git a/cmd/launcher/web/web_test.go b/cmd/launcher/web/web_test.go index 1078f1a..d5e43b0 100644 --- a/cmd/launcher/web/web_test.go +++ b/cmd/launcher/web/web_test.go @@ -19,6 +19,7 @@ import ( "github.com/sylumi/agentkit/model" "github.com/sylumi/agentkit/model/openaimodel" "github.com/sylumi/agentkit/session" + "github.com/sylumi/agentkit/session/inmemory" webui "github.com/sylumi/agentkit/web" ) @@ -160,7 +161,7 @@ func TestWebUsesConfiguredServiceWithoutUI(t *testing.T) { yield(e, nil) } }} - store := session.InMemoryService() + store := inmemory.New() cfg := launcher.Config{Agent: a, SessionService: store, AppName: "custom-app", UserID: "alice"} h, err := NewHandler(cfg, nil) if err != nil { diff --git a/runner/runner_test.go b/runner/runner_test.go index 61e042d..e63f908 100644 --- a/runner/runner_test.go +++ b/runner/runner_test.go @@ -16,6 +16,7 @@ import ( "github.com/sylumi/agentkit/model" "github.com/sylumi/agentkit/runner" "github.com/sylumi/agentkit/session" + "github.com/sylumi/agentkit/session/inmemory" "github.com/sylumi/agentkit/tool" "github.com/sylumi/agentkit/tool/functiontool" ) @@ -56,7 +57,7 @@ func (s serviceStub) AppendEvent(ctx context.Context, current session.Session, e func createSession(t *testing.T) session.Service { t.Helper() - s := session.InMemoryService() + s := inmemory.New() if _, err := s.Create(t.Context(), &session.CreateRequest{AppName: "app", UserID: "user", SessionID: "session"}); err != nil { t.Fatal(err) } diff --git a/server/server_test.go b/server/server_test.go index ee57485..7dd7614 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -19,6 +19,7 @@ import ( "github.com/sylumi/agentkit/model" "github.com/sylumi/agentkit/server" "github.com/sylumi/agentkit/session" + "github.com/sylumi/agentkit/session/inmemory" "github.com/sylumi/agentkit/tool" "github.com/sylumi/agentkit/tool/functiontool" ) @@ -43,7 +44,7 @@ func echoAgent(ctx context.Context, inv *agent.InvocationContext) iter.Seq2[*ses func newServer(t *testing.T, a agent.Agent) (*server.Server, session.Service) { t.Helper() - store := session.InMemoryService() + store := inmemory.New() s, err := server.New(server.Config{AppName: "test", Agent: a, SessionService: store}) if err != nil { t.Fatal(err) @@ -88,7 +89,7 @@ func create(t *testing.T, h http.Handler) string { } func TestConfigAndHTTPValidation(t *testing.T) { - base := server.Config{AppName: "test", Agent: agentFunc(echoAgent), SessionService: session.InMemoryService()} + base := server.Config{AppName: "test", Agent: agentFunc(echoAgent), SessionService: inmemory.New()} for _, mutate := range []func(*server.Config){ func(c *server.Config) { c.AppName = "" }, func(c *server.Config) { c.Agent = nil }, func(c *server.Config) { c.SessionService = nil }, @@ -497,7 +498,7 @@ func TestSessionSnapshotDuringRunCompletion(t *testing.T) { readContext, cancel := context.WithCancel(t.Context()) defer cancel() store := &snapshotStore{ - Service: session.InMemoryService(), readContext: readContext, + Service: inmemory.New(), readContext: readContext, read: make(chan struct{}), resume: make(chan struct{}), } started, finish, committed := make(chan struct{}), make(chan struct{}), make(chan struct{}) diff --git a/session/inmemory.go b/session/inmemory.go deleted file mode 100644 index 6624f83..0000000 --- a/session/inmemory.go +++ /dev/null @@ -1,193 +0,0 @@ -package session - -import ( - "context" - "fmt" - "iter" - "slices" - "strings" - "sync" - "time" - - "github.com/google/uuid" -) - -type inMemoryService struct { - mu sync.RWMutex - sessions map[sessionKey]*session -} - -// InMemoryService creates a service whose data lasts for the life of the instance. -func InMemoryService() Service { - return &inMemoryService{sessions: make(map[sessionKey]*session)} -} - -func (s *inMemoryService) Create(ctx context.Context, req *CreateRequest) (*CreateResponse, error) { - if strings.TrimSpace(req.AppName) == "" || strings.TrimSpace(req.UserID) == "" { - return nil, fmt.Errorf("session: app_name and user_id are required") - } - - sessionID := req.SessionID - if sessionID == "" { - sessionID = uuid.NewString() - } - key := sessionKey{ - appName: req.AppName, - userID: req.UserID, - sessionID: sessionID, - } - s.mu.Lock() - defer s.mu.Unlock() - if _, exists := s.sessions[key]; exists { - return nil, fmt.Errorf("%w: %q", ErrAlreadyExists, sessionID) - } - record := &session{ - key: key, - updatedAt: time.Now().UTC(), - } - s.sessions[key] = record - return &CreateResponse{Session: copySession(record, true)}, nil -} - -func (s *inMemoryService) Get(ctx context.Context, req *GetRequest) (*GetResponse, error) { - if strings.TrimSpace(req.AppName) == "" || strings.TrimSpace(req.UserID) == "" || strings.TrimSpace(req.SessionID) == "" { - return nil, fmt.Errorf("session: app_name, user_id and session_id are required") - } - key := sessionKey{ - appName: req.AppName, - userID: req.UserID, - sessionID: req.SessionID, - } - s.mu.RLock() - defer s.mu.RUnlock() - record, ok := s.sessions[key] - if !ok { - return nil, fmt.Errorf("%w: %q", ErrNotFound, req.SessionID) - } - return &GetResponse{Session: copySession(record, true)}, nil -} - -func (s *inMemoryService) List(ctx context.Context, req *ListRequest) (*ListResponse, error) { - if strings.TrimSpace(req.AppName) == "" || strings.TrimSpace(req.UserID) == "" { - return nil, fmt.Errorf("session: app_name and user_id are required") - } - s.mu.RLock() - defer s.mu.RUnlock() - result := &ListResponse{Sessions: []Session{}} - for key, record := range s.sessions { - if key.appName == req.AppName && key.userID == req.UserID { - result.Sessions = append(result.Sessions, copySession(record, false)) - } - } - slices.SortFunc(result.Sessions, func(a, b Session) int { - return strings.Compare(a.ID(), b.ID()) - }) - return result, nil -} - -func (s *inMemoryService) Delete(ctx context.Context, req *DeleteRequest) error { - if strings.TrimSpace(req.AppName) == "" || strings.TrimSpace(req.UserID) == "" || strings.TrimSpace(req.SessionID) == "" { - return fmt.Errorf("session: app_name, user_id and session_id are required") - } - key := sessionKey{ - appName: req.AppName, - userID: req.UserID, - sessionID: req.SessionID, - } - s.mu.Lock() - defer s.mu.Unlock() - delete(s.sessions, key) - return nil -} - -func (s *inMemoryService) AppendEvent(ctx context.Context, current Session, event *Event) error { - view, ok := current.(*session) - if !ok || view == nil { - return ErrInvalidSession - } - if event == nil { - return fmt.Errorf("session: event is required") - } - if event.Partial { - return nil - } - storedEvent := *event - - s.mu.Lock() - defer s.mu.Unlock() - view.mu.Lock() - defer view.mu.Unlock() - record, ok := s.sessions[view.key] - if !ok { - return fmt.Errorf("%w: %q", ErrNotFound, view.ID()) - } - record.events = append(record.events, &storedEvent) - record.updatedAt = time.Now().UTC() - view.events = append(view.events, &storedEvent) - view.updatedAt = record.updatedAt - return nil -} - -// copySession is called while the service lock is held. -func copySession(record *session, includeEvents bool) *session { - view := &session{ - key: record.key, - updatedAt: record.updatedAt, - } - if includeEvents { - view.events = slices.Clone(record.events) - } - return view -} - -type sessionKey struct { - appName string - userID string - sessionID string -} - -type session struct { - key sessionKey - - mu sync.RWMutex - events []*Event - updatedAt time.Time -} - -func (s *session) ID() string { return s.key.sessionID } -func (s *session) AppName() string { return s.key.appName } -func (s *session) UserID() string { return s.key.userID } - -func (s *session) Events() Events { - s.mu.RLock() - defer s.mu.RUnlock() - return events(s.events) -} - -// LastUpdateTime reports the last service commit time, not the event timestamp. -func (s *session) LastUpdateTime() time.Time { - s.mu.RLock() - defer s.mu.RUnlock() - return s.updatedAt -} - -type events []*Event - -func (e events) All() iter.Seq[*Event] { - return func(yield func(*Event) bool) { - for _, event := range e { - if !yield(event) { - return - } - } - } -} - -func (e events) Len() int { return len(e) } - -func (e events) At(index int) *Event { - if index < 0 || index >= len(e) { - return nil - } - return e[index] -} diff --git a/session/inmemory/service.go b/session/inmemory/service.go new file mode 100644 index 0000000..9e22cae --- /dev/null +++ b/session/inmemory/service.go @@ -0,0 +1,131 @@ +// Package inmemory provides an in-memory session service. +package inmemory + +import ( + "context" + "fmt" + "slices" + "strings" + "sync" + "time" + + "github.com/google/uuid" + + "github.com/sylumi/agentkit/session" +) + +type service struct { + mu sync.RWMutex + sessions map[sessionKey]*localSession +} + +// New creates a service whose data lasts for the life of the instance. +func New() session.Service { + return &service{sessions: make(map[sessionKey]*localSession)} +} + +func (s *service) Create(ctx context.Context, req *session.CreateRequest) (*session.CreateResponse, error) { + if strings.TrimSpace(req.AppName) == "" || strings.TrimSpace(req.UserID) == "" { + return nil, fmt.Errorf("session: app_name and user_id are required") + } + + sessionID := req.SessionID + if sessionID == "" { + sessionID = uuid.NewString() + } + key := sessionKey{ + appName: req.AppName, + userID: req.UserID, + sessionID: sessionID, + } + s.mu.Lock() + defer s.mu.Unlock() + if _, exists := s.sessions[key]; exists { + return nil, fmt.Errorf("%w: %q", session.ErrAlreadyExists, sessionID) + } + record := &localSession{ + key: key, + updatedAt: time.Now().UTC(), + } + s.sessions[key] = record + return &session.CreateResponse{Session: copySession(record, true)}, nil +} + +func (s *service) Get(ctx context.Context, req *session.GetRequest) (*session.GetResponse, error) { + if strings.TrimSpace(req.AppName) == "" || strings.TrimSpace(req.UserID) == "" || strings.TrimSpace(req.SessionID) == "" { + return nil, fmt.Errorf("session: app_name, user_id and session_id are required") + } + key := sessionKey{ + appName: req.AppName, + userID: req.UserID, + sessionID: req.SessionID, + } + s.mu.RLock() + defer s.mu.RUnlock() + record, ok := s.sessions[key] + if !ok { + return nil, fmt.Errorf("%w: %q", session.ErrNotFound, req.SessionID) + } + return &session.GetResponse{Session: copySession(record, true)}, nil +} + +func (s *service) List(ctx context.Context, req *session.ListRequest) (*session.ListResponse, error) { + if strings.TrimSpace(req.AppName) == "" || strings.TrimSpace(req.UserID) == "" { + return nil, fmt.Errorf("session: app_name and user_id are required") + } + s.mu.RLock() + defer s.mu.RUnlock() + result := &session.ListResponse{Sessions: []session.Session{}} + for key, record := range s.sessions { + if key.appName == req.AppName && key.userID == req.UserID { + result.Sessions = append(result.Sessions, copySession(record, false)) + } + } + slices.SortFunc(result.Sessions, func(a, b session.Session) int { + return strings.Compare(a.ID(), b.ID()) + }) + return result, nil +} + +func (s *service) Delete(ctx context.Context, req *session.DeleteRequest) error { + if strings.TrimSpace(req.AppName) == "" || strings.TrimSpace(req.UserID) == "" || strings.TrimSpace(req.SessionID) == "" { + return fmt.Errorf("session: app_name, user_id and session_id are required") + } + key := sessionKey{ + appName: req.AppName, + userID: req.UserID, + sessionID: req.SessionID, + } + s.mu.Lock() + defer s.mu.Unlock() + delete(s.sessions, key) + return nil +} + +func (s *service) AppendEvent(ctx context.Context, current session.Session, event *session.Event) error { + view, ok := current.(*localSession) + if !ok || view == nil { + return session.ErrInvalidSession + } + if event == nil { + return fmt.Errorf("session: event is required") + } + if event.Partial { + return nil + } + storedEvent := *event + + s.mu.Lock() + defer s.mu.Unlock() + view.mu.Lock() + defer view.mu.Unlock() + record, ok := s.sessions[view.key] + if !ok { + return fmt.Errorf("%w: %q", session.ErrNotFound, view.ID()) + } + record.events = append(record.events, &storedEvent) + record.updatedAt = time.Now().UTC() + view.events = append(view.events, &storedEvent) + view.updatedAt = record.updatedAt + return nil +} diff --git a/session/inmemory_test.go b/session/inmemory/service_test.go similarity index 76% rename from session/inmemory_test.go rename to session/inmemory/service_test.go index 9856dca..85ab9bd 100644 --- a/session/inmemory_test.go +++ b/session/inmemory/service_test.go @@ -1,4 +1,4 @@ -package session_test +package inmemory_test import ( "bytes" @@ -14,11 +14,12 @@ import ( "github.com/sylumi/agentkit/model" "github.com/sylumi/agentkit/session" + "github.com/sylumi/agentkit/session/inmemory" ) func createSession(t *testing.T, id string) (session.Service, session.Session) { t.Helper() - service := session.InMemoryService() + service := inmemory.New() created, err := service.Create(t.Context(), &session.CreateRequest{AppName: "app", UserID: "user", SessionID: id}) if err != nil { t.Fatal(err) @@ -36,7 +37,7 @@ func getSession(t *testing.T, service session.Service, view session.Session) ses } func TestServiceLifecycleAndIdentityIsolation(t *testing.T) { - service := session.InMemoryService() + service := inmemory.New() ctx := t.Context() request := &session.CreateRequest{AppName: "app", UserID: "user"} generated, err := service.Create(ctx, request) @@ -226,7 +227,7 @@ func TestAppendRejectsInvalidAndMissingSessions(t *testing.T) { } func TestServiceInputValidation(t *testing.T) { - service := session.InMemoryService() + service := inmemory.New() _, createErr := service.Create(t.Context(), &session.CreateRequest{AppName: " ", UserID: "user"}) _, getErr := service.Get(t.Context(), &session.GetRequest{AppName: "app", UserID: "user"}) _, listErr := service.List(t.Context(), &session.ListRequest{AppName: "app"}) @@ -280,3 +281,88 @@ func TestServiceConcurrentAccess(t *testing.T) { } } } + +func userEvent(text string) *session.Event { + event := session.NewEvent("invocation") + event.Author = "user" + event.Message = &model.Message{Role: model.RoleUser, Parts: []model.Part{model.NewTextPart(text)}} + return event +} + +func toolCallMessage() *model.Message { + return &model.Message{Role: model.RoleAssistant, Parts: []model.Part{{ + Kind: model.PartToolCall, + ToolCall: &model.ToolCallPart{ + ID: "call-1", Name: "lookup", Arguments: json.RawMessage(`{ "id": 9007199254740993 }`), + }, + }}} +} + +func TestAppendEventPreservesPayloadWithoutValidation(t *testing.T) { + for _, tc := range []struct { + name string + event *session.Event + }{ + {"empty record", &session.Event{}}, + {"message without parts", &session.Event{Message: &model.Message{Role: model.RoleUser}}}, + {"metadata without generation", &session.Event{Metadata: &model.ResponseMetadata{ResponseID: "response"}}}, + } { + t.Run(tc.name, func(t *testing.T) { + service, view := createSession(t, "payload") + if err := service.AppendEvent(t.Context(), view, tc.event); err != nil { + t.Fatal(err) + } + if got := getSession(t, service, view).Events().At(0); !reflect.DeepEqual(got, tc.event) { + t.Fatalf("event changed: got %+v, want %+v", got, tc.event) + } + }) + } +} + +func TestAppendEventPreservesInput(t *testing.T) { + service, view := createSession(t, "json") + event := session.NewEvent("invocation") + event.Author = "assistant" + event.Message = &model.Message{Role: model.RoleAssistant, Parts: []model.Part{model.NewTextPart("answer")}} + event.StopReason = model.StopReasonStop + count := int64(0) + event.Usage = &model.Usage{InputTokens: &count} + event.Metadata = &model.ResponseMetadata{Provider: "test", ResponseID: "provider-response"} + + before, err := json.Marshal(event) + if err != nil { + t.Fatal(err) + } + if err := service.AppendEvent(t.Context(), view, event); err != nil { + t.Fatal(err) + } + after, err := json.Marshal(event) + if err != nil || !bytes.Equal(before, after) { + t.Fatalf("append changed event: %s, %v", after, err) + } + if got := getSession(t, service, view).Events().At(0); !reflect.DeepEqual(got, event) { + t.Fatalf("stored event changed: got %+v, want %+v", got, event) + } +} + +func TestPartialEventsDoNotPersist(t *testing.T) { + service, view := createSession(t, "stream") + updated := view.LastUpdateTime() + for _, delta := range []model.Event{ + model.PartStart{Index: 0, Kind: model.PartThinking, ThinkingKind: model.ThinkingSummary}, + model.ThinkingDelta{Index: 0, Delta: "Checking"}, + model.TextDelta{Index: 1, Delta: "Hello"}, + model.ToolCallDelta{Index: 2, ID: "call", Name: "lookup", Arguments: `{"city":`}, + model.PartEnd{Index: 0}, + } { + event := session.NewEvent("run") + event.Author, event.Partial, event.Delta = "agent", true, delta + if err := service.AppendEvent(t.Context(), view, event); err != nil { + t.Fatal(err) + } + } + stored := getSession(t, service, view) + if view.Events().Len() != 0 || stored.Events().Len() != 0 || !view.LastUpdateTime().Equal(updated) || !stored.LastUpdateTime().Equal(updated) { + t.Fatal("partial events changed session history or update time") + } +} diff --git a/session/inmemory/session.go b/session/inmemory/session.go new file mode 100644 index 0000000..9ecc909 --- /dev/null +++ b/session/inmemory/session.go @@ -0,0 +1,74 @@ +package inmemory + +import ( + "iter" + "slices" + "sync" + "time" + + "github.com/sylumi/agentkit/session" +) + +// copySession is called while the service lock is held. +func copySession(record *localSession, includeEvents bool) *localSession { + view := &localSession{ + key: record.key, + updatedAt: record.updatedAt, + } + if includeEvents { + view.events = slices.Clone(record.events) + } + return view +} + +type sessionKey struct { + appName string + userID string + sessionID string +} + +type localSession struct { + key sessionKey + + mu sync.RWMutex + events []*session.Event + updatedAt time.Time +} + +func (s *localSession) ID() string { return s.key.sessionID } +func (s *localSession) AppName() string { return s.key.appName } +func (s *localSession) UserID() string { return s.key.userID } + +func (s *localSession) Events() session.Events { + s.mu.RLock() + defer s.mu.RUnlock() + return events(s.events) +} + +// LastUpdateTime reports the last service commit time, not the event timestamp. +func (s *localSession) LastUpdateTime() time.Time { + s.mu.RLock() + defer s.mu.RUnlock() + return s.updatedAt +} + +type events []*session.Event + +func (e events) All() iter.Seq[*session.Event] { + return func(yield func(*session.Event) bool) { + for _, event := range e { + if !yield(event) { + return + } + } + } +} + +func (e events) Len() int { return len(e) } + +func (e events) At(index int) *session.Event { + if index < 0 || index >= len(e) { + return nil + } + return e[index] +} diff --git a/session/session_test.go b/session/session_test.go index 2c21b7d..02bed52 100644 --- a/session/session_test.go +++ b/session/session_test.go @@ -1,7 +1,6 @@ package session_test import ( - "bytes" "encoding/json" "reflect" "testing" @@ -10,46 +9,8 @@ import ( "github.com/sylumi/agentkit/session" ) -func userEvent(text string) *session.Event { +func TestEventJSONPreservesContent(t *testing.T) { event := session.NewEvent("invocation") - event.Author = "user" - event.Message = &model.Message{Role: model.RoleUser, Parts: []model.Part{model.NewTextPart(text)}} - return event -} - -func toolCallMessage() *model.Message { - return &model.Message{Role: model.RoleAssistant, Parts: []model.Part{{ - Kind: model.PartToolCall, - ToolCall: &model.ToolCallPart{ - ID: "call-1", Name: "lookup", Arguments: json.RawMessage(`{ "id": 9007199254740993 }`), - }, - }}} -} - -func TestAppendEventPreservesPayloadWithoutValidation(t *testing.T) { - for _, tc := range []struct { - name string - event *session.Event - }{ - {"empty record", &session.Event{}}, - {"message without parts", &session.Event{Message: &model.Message{Role: model.RoleUser}}}, - {"metadata without generation", &session.Event{Metadata: &model.ResponseMetadata{ResponseID: "response"}}}, - } { - t.Run(tc.name, func(t *testing.T) { - service, view := createSession(t, "payload") - if err := service.AppendEvent(t.Context(), view, tc.event); err != nil { - t.Fatal(err) - } - if got := getSession(t, service, view).Events().At(0); !reflect.DeepEqual(got, tc.event) { - t.Fatalf("event changed: got %+v, want %+v", got, tc.event) - } - }) - } -} - -func TestEventJSONAndAppendPreserveContent(t *testing.T) { - service, view := createSession(t, "json") - event := userEvent("hello") event.Author = "assistant" event.Message = &model.Message{Role: model.RoleAssistant, Parts: []model.Part{model.NewTextPart("answer")}} event.StopReason = model.StopReasonStop @@ -61,13 +22,6 @@ func TestEventJSONAndAppendPreserveContent(t *testing.T) { if err != nil { t.Fatal(err) } - if err := service.AppendEvent(t.Context(), view, event); err != nil { - t.Fatal(err) - } - after, err := json.Marshal(event) - if err != nil || !bytes.Equal(before, after) { - t.Fatalf("append changed event: %s, %v", after, err) - } var decoded session.Event if err := json.Unmarshal(before, &decoded); err != nil { t.Fatal(err) @@ -80,9 +34,7 @@ func TestEventJSONAndAppendPreserveContent(t *testing.T) { } } -func TestPartialEventsRoundTripWithoutPersistence(t *testing.T) { - service, view := createSession(t, "stream") - updated := view.LastUpdateTime() +func TestPartialEventsRoundTrip(t *testing.T) { for _, delta := range []model.Event{ model.PartStart{Index: 0, Kind: model.PartThinking, ThinkingKind: model.ThinkingSummary}, model.ThinkingDelta{Index: 0, Delta: "Checking"}, @@ -100,12 +52,5 @@ func TestPartialEventsRoundTripWithoutPersistence(t *testing.T) { if err := json.Unmarshal(data, &decoded); err != nil || !reflect.DeepEqual(event, &decoded) { t.Fatalf("partial event changed during JSON round trip: %s, %v", data, err) } - if err := service.AppendEvent(t.Context(), view, &decoded); err != nil { - t.Fatal(err) - } - } - stored := getSession(t, service, view) - if view.Events().Len() != 0 || stored.Events().Len() != 0 || !view.LastUpdateTime().Equal(updated) || !stored.LastUpdateTime().Equal(updated) { - t.Fatal("partial events changed session history or update time") } }