diff --git a/internal/connector/connector.go b/internal/connector/connector.go index 04170efd6..a23ae121c 100644 --- a/internal/connector/connector.go +++ b/internal/connector/connector.go @@ -68,23 +68,20 @@ func (c *connectorClient) SetSessionStoreSetter(sessionStoreSetter sessions.SetS c.sessionStoreSetter = sessionStoreSetter } -func (c *connectorClient) SetSessionStore(ctx context.Context, store sessions.SessionStore) { +func (c *connectorClient) SetSessionStore(ctx context.Context, syncID string, store sessions.SessionStore) { if c.sessionStoreSetter == nil { - // Demoted from Warn to Debug: this path is the normal case for - // any connector that didn't opt into session storage (i.e., - // wrapper.run never set cw.SessionServer to non-nil, so - // SetSessionStoreSetter received nil at startup). The syncer - // still calls SetSessionStore unconditionally on every sync, so - // at Warn level this fires once per sync per non-session-store - // connector — pure noise that masked real warnings downstream. - // Debug keeps the "you forgot to wire this" signal available for - // developers running at debug level without polluting prod logs. - // See https://github.com/ConductorOne/baton-sdk/issues/907. l := ctxzap.Extract(ctx) l.Debug("connectorClient's session store is nil — connector did not opt into session storage") return } - c.sessionStoreSetter.SetSessionStore(ctx, store) + c.sessionStoreSetter.SetSessionStore(ctx, syncID, store) +} + +func (c *connectorClient) RemoveSessionStore(ctx context.Context, syncID string) { + if c.sessionStoreSetter == nil { + return + } + c.sessionStoreSetter.RemoveSessionStore(ctx, syncID) } var ErrConnectorNotImplemented = errors.New("client does not implement connector connectorV2") diff --git a/pkg/session/session_server.go b/pkg/session/session_server.go index 622ec1eeb..fb0c1072e 100644 --- a/pkg/session/session_server.go +++ b/pkg/session/session_server.go @@ -5,6 +5,7 @@ import ( "fmt" "log" "net" + "sync" v1 "github.com/conductorone/baton-sdk/pb/c1/connectorapi/baton/v1" "github.com/conductorone/baton-sdk/pkg/types/sessions" @@ -16,35 +17,50 @@ var _ v1.BatonSessionServiceServer = (*GRPCSessionServer)(nil) type GRPCSessionServer struct { // v1.UnimplementedBatonSessionServiceServer - store sessions.SessionStore + mu sync.RWMutex + stores map[string]sessions.SessionStore } func NewGRPCSessionServer() *GRPCSessionServer { - return &GRPCSessionServer{} + return &GRPCSessionServer{ + stores: make(map[string]sessions.SessionStore), + } } type SetSessionStore interface { - SetSessionStore(ctx context.Context, store sessions.SessionStore) + SetSessionStore(ctx context.Context, syncID string, store sessions.SessionStore) + RemoveSessionStore(ctx context.Context, syncID string) } -func (s *GRPCSessionServer) SetSessionStore(ctx context.Context, store sessions.SessionStore) { - s.store = store +func (s *GRPCSessionServer) SetSessionStore(ctx context.Context, syncID string, store sessions.SessionStore) { + s.mu.Lock() + defer s.mu.Unlock() + s.stores[syncID] = store } -func (s *GRPCSessionServer) Validate() error { - if s.store == nil { - return fmt.Errorf("session store is not set") - } +func (s *GRPCSessionServer) RemoveSessionStore(ctx context.Context, syncID string) { + s.mu.Lock() + defer s.mu.Unlock() + delete(s.stores, syncID) +} - return nil +func (s *GRPCSessionServer) getStore(syncID string) (sessions.SessionStore, error) { + s.mu.RLock() + defer s.mu.RUnlock() + store, ok := s.stores[syncID] + if !ok { + return nil, fmt.Errorf("session store not found for sync_id %q", syncID) + } + return store, nil } func (s *GRPCSessionServer) Get(ctx context.Context, req *v1.GetRequest) (*v1.GetResponse, error) { - if err := s.Validate(); err != nil { + store, err := s.getStore(req.GetSyncId()) + if err != nil { return nil, err } - value, found, err := s.store.Get(ctx, req.GetKey(), sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) + value, found, err := store.Get(ctx, req.GetKey(), sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) if err != nil { return nil, fmt.Errorf("failed to get value from cache: %w", err) } @@ -56,11 +72,12 @@ func (s *GRPCSessionServer) Get(ctx context.Context, req *v1.GetRequest) (*v1.Ge } func (s *GRPCSessionServer) GetMany(ctx context.Context, req *v1.GetManyRequest) (*v1.GetManyResponse, error) { - if err := s.Validate(); err != nil { + store, err := s.getStore(req.GetSyncId()) + if err != nil { return nil, err } - values, unprocessedKeys, err := s.store.GetMany( + values, unprocessedKeys, err := store.GetMany( ctx, req.GetKeys(), sessions.WithSyncID(req.GetSyncId()), @@ -86,11 +103,12 @@ func (s *GRPCSessionServer) GetMany(ctx context.Context, req *v1.GetManyRequest) } func (s *GRPCSessionServer) Set(ctx context.Context, req *v1.SetRequest) (*v1.SetResponse, error) { - if err := s.Validate(); err != nil { + store, err := s.getStore(req.GetSyncId()) + if err != nil { return nil, err } - err := s.store.Set(ctx, req.GetKey(), req.GetValue(), sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) + err = store.Set(ctx, req.GetKey(), req.GetValue(), sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) if err != nil { return nil, fmt.Errorf("failed to set value in cache: %w", err) } @@ -99,11 +117,12 @@ func (s *GRPCSessionServer) Set(ctx context.Context, req *v1.SetRequest) (*v1.Se } func (s *GRPCSessionServer) SetMany(ctx context.Context, req *v1.SetManyRequest) (*v1.SetManyResponse, error) { - if err := s.Validate(); err != nil { + store, err := s.getStore(req.GetSyncId()) + if err != nil { return nil, err } - err := s.store.SetMany(ctx, req.GetValues(), sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) + err = store.SetMany(ctx, req.GetValues(), sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) if err != nil { return nil, fmt.Errorf("failed to set many values in cache: %w", err) } @@ -112,11 +131,12 @@ func (s *GRPCSessionServer) SetMany(ctx context.Context, req *v1.SetManyRequest) } func (s *GRPCSessionServer) Delete(ctx context.Context, req *v1.DeleteRequest) (*v1.DeleteResponse, error) { - if err := s.Validate(); err != nil { + store, err := s.getStore(req.GetSyncId()) + if err != nil { return nil, err } - err := s.store.Delete(ctx, req.GetKey(), sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) + err = store.Delete(ctx, req.GetKey(), sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) if err != nil { return nil, fmt.Errorf("failed to delete value from cache: %w", err) } @@ -125,12 +145,13 @@ func (s *GRPCSessionServer) Delete(ctx context.Context, req *v1.DeleteRequest) ( } func (s *GRPCSessionServer) DeleteMany(ctx context.Context, req *v1.DeleteManyRequest) (*v1.DeleteManyResponse, error) { - if err := s.Validate(); err != nil { + store, err := s.getStore(req.GetSyncId()) + if err != nil { return nil, err } for _, key := range req.GetKeys() { - err := s.store.Delete( + err := store.Delete( ctx, key, sessions.WithSyncID(req.GetSyncId()), @@ -145,13 +166,13 @@ func (s *GRPCSessionServer) DeleteMany(ctx context.Context, req *v1.DeleteManyRe } func (s *GRPCSessionServer) Clear(ctx context.Context, req *v1.ClearRequest) (*v1.ClearResponse, error) { - if s.store == nil { - // we sometimes clean up the session store after the connector is done + store, err := s.getStore(req.GetSyncId()) + if err != nil { ctxzap.Extract(ctx).Warn("session store is not set") return &v1.ClearResponse{}, nil } - err := s.store.Clear(ctx, sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) + err = store.Clear(ctx, sessions.WithSyncID(req.GetSyncId()), sessions.WithPrefix(req.GetPrefix())) if err != nil { return nil, fmt.Errorf("failed to clear cache: %w", err) } @@ -160,11 +181,12 @@ func (s *GRPCSessionServer) Clear(ctx context.Context, req *v1.ClearRequest) (*v } func (s *GRPCSessionServer) GetAll(ctx context.Context, req *v1.GetAllRequest) (*v1.GetAllResponse, error) { - if err := s.Validate(); err != nil { + store, err := s.getStore(req.GetSyncId()) + if err != nil { return nil, err } - values, nextPageToken, err := s.store.GetAll( + values, nextPageToken, err := store.GetAll( ctx, req.PageToken, sessions.WithSyncID(req.GetSyncId()), diff --git a/pkg/session/session_server_test.go b/pkg/session/session_server_test.go new file mode 100644 index 000000000..89c02fc77 --- /dev/null +++ b/pkg/session/session_server_test.go @@ -0,0 +1,269 @@ +package session + +import ( + "context" + "fmt" + "sync" + "testing" + + v1 "github.com/conductorone/baton-sdk/pb/c1/connectorapi/baton/v1" + "github.com/conductorone/baton-sdk/pkg/types/sessions" + "github.com/stretchr/testify/require" +) + +func TestGRPCSessionServer_ConcurrentSyncs(t *testing.T) { + server := NewGRPCSessionServer() + ctx := context.Background() + + syncA := "sync-A" + syncB := "sync-B" + + storeA := &inMemorySessionStore{data: make(map[storeKey][]byte)} + storeB := &inMemorySessionStore{data: make(map[storeKey][]byte)} + + server.SetSessionStore(ctx, syncA, storeA) + server.SetSessionStore(ctx, syncB, storeB) + + _, err := server.Set(ctx, v1.SetRequest_builder{ + SyncId: syncA, + Key: "key1", + Value: []byte("value-A"), + }.Build()) + require.NoError(t, err) + + _, err = server.Set(ctx, v1.SetRequest_builder{ + SyncId: syncB, + Key: "key1", + Value: []byte("value-B"), + }.Build()) + require.NoError(t, err) + + respA, err := server.Get(ctx, v1.GetRequest_builder{ + SyncId: syncA, + Key: "key1", + }.Build()) + require.NoError(t, err) + require.True(t, respA.GetFound()) + require.Equal(t, []byte("value-A"), respA.GetValue()) + + respB, err := server.Get(ctx, v1.GetRequest_builder{ + SyncId: syncB, + Key: "key1", + }.Build()) + require.NoError(t, err) + require.True(t, respB.GetFound()) + require.Equal(t, []byte("value-B"), respB.GetValue()) + + // Sync A should not see sync B's data and vice versa. + respA2, err := server.Get(ctx, v1.GetRequest_builder{ + SyncId: syncA, + Key: "key1", + }.Build()) + require.NoError(t, err) + require.Equal(t, []byte("value-A"), respA2.GetValue(), "sync A read sync B's data") +} + +func TestGRPCSessionServer_ConcurrentReadWrite(t *testing.T) { + server := NewGRPCSessionServer() + ctx := context.Background() + + const numSyncs = 10 + const numOps = 50 + + stores := make([]*inMemorySessionStore, numSyncs) + for i := range numSyncs { + stores[i] = &inMemorySessionStore{data: make(map[storeKey][]byte)} + server.SetSessionStore(ctx, fmt.Sprintf("sync-%d", i), stores[i]) + } + + var wg sync.WaitGroup + for i := range numSyncs { + wg.Add(1) + go func(syncIdx int) { + defer wg.Done() + syncID := fmt.Sprintf("sync-%d", syncIdx) + for j := range numOps { + key := fmt.Sprintf("key-%d", j) + value := []byte(fmt.Sprintf("value-%d-%d", syncIdx, j)) + + _, err := server.Set(ctx, v1.SetRequest_builder{ + SyncId: syncID, + Key: key, + Value: value, + }.Build()) + require.NoError(t, err) + + resp, err := server.Get(ctx, v1.GetRequest_builder{ + SyncId: syncID, + Key: key, + }.Build()) + require.NoError(t, err) + require.True(t, resp.GetFound()) + require.Equal(t, value, resp.GetValue(), + "sync %d read wrong value for %s", syncIdx, key) + } + }(i) + } + wg.Wait() +} + +func TestGRPCSessionServer_RemoveSessionStore(t *testing.T) { + server := NewGRPCSessionServer() + ctx := context.Background() + + store := &inMemorySessionStore{data: make(map[storeKey][]byte)} + server.SetSessionStore(ctx, "sync-1", store) + + _, err := server.Set(ctx, v1.SetRequest_builder{ + SyncId: "sync-1", + Key: "key1", + Value: []byte("value1"), + }.Build()) + require.NoError(t, err) + + server.RemoveSessionStore(ctx, "sync-1") + + _, err = server.Get(ctx, v1.GetRequest_builder{ + SyncId: "sync-1", + Key: "key1", + }.Build()) + require.Error(t, err) + require.Contains(t, err.Error(), "session store not found") +} + +func TestGRPCSessionServer_UnregisteredSyncID(t *testing.T) { + server := NewGRPCSessionServer() + ctx := context.Background() + + _, err := server.Get(ctx, v1.GetRequest_builder{ + SyncId: "nonexistent", + Key: "key1", + }.Build()) + require.Error(t, err) + require.Contains(t, err.Error(), "session store not found") +} + +func TestGRPCSessionServer_ClearUnregistered(t *testing.T) { + server := NewGRPCSessionServer() + ctx := context.Background() + + resp, err := server.Clear(ctx, v1.ClearRequest_builder{ + SyncId: "nonexistent", + }.Build()) + require.NoError(t, err) + require.NotNil(t, resp) +} + +// inMemorySessionStore is a minimal in-memory session store for testing. +type storeKey struct { + syncID string + key string +} + +type inMemorySessionStore struct { + mu sync.RWMutex + data map[storeKey][]byte +} + +func (s *inMemorySessionStore) Get(ctx context.Context, key string, opt ...sessions.SessionStoreOption) ([]byte, bool, error) { + bag, err := applyOpts(ctx, opt...) + if err != nil { + return nil, false, err + } + s.mu.RLock() + defer s.mu.RUnlock() + v, ok := s.data[storeKey{syncID: bag.SyncID, key: bag.Prefix + key}] + return v, ok, nil +} + +func (s *inMemorySessionStore) GetMany(ctx context.Context, keys []string, opt ...sessions.SessionStoreOption) (map[string][]byte, []string, error) { + bag, err := applyOpts(ctx, opt...) + if err != nil { + return nil, nil, err + } + s.mu.RLock() + defer s.mu.RUnlock() + result := make(map[string][]byte) + for _, k := range keys { + if v, ok := s.data[storeKey{syncID: bag.SyncID, key: bag.Prefix + k}]; ok { + result[k] = v + } + } + return result, nil, nil +} + +func (s *inMemorySessionStore) Set(ctx context.Context, key string, value []byte, opt ...sessions.SessionStoreOption) error { + bag, err := applyOpts(ctx, opt...) + if err != nil { + return err + } + s.mu.Lock() + defer s.mu.Unlock() + s.data[storeKey{syncID: bag.SyncID, key: bag.Prefix + key}] = value + return nil +} + +func (s *inMemorySessionStore) SetMany(ctx context.Context, values map[string][]byte, opt ...sessions.SessionStoreOption) error { + bag, err := applyOpts(ctx, opt...) + if err != nil { + return err + } + s.mu.Lock() + defer s.mu.Unlock() + for k, v := range values { + s.data[storeKey{syncID: bag.SyncID, key: bag.Prefix + k}] = v + } + return nil +} + +func (s *inMemorySessionStore) Delete(ctx context.Context, key string, opt ...sessions.SessionStoreOption) error { + bag, err := applyOpts(ctx, opt...) + if err != nil { + return err + } + s.mu.Lock() + defer s.mu.Unlock() + delete(s.data, storeKey{syncID: bag.SyncID, key: bag.Prefix + key}) + return nil +} + +func (s *inMemorySessionStore) Clear(ctx context.Context, opt ...sessions.SessionStoreOption) error { + bag, err := applyOpts(ctx, opt...) + if err != nil { + return err + } + s.mu.Lock() + defer s.mu.Unlock() + for k := range s.data { + if k.syncID == bag.SyncID { + delete(s.data, k) + } + } + return nil +} + +func (s *inMemorySessionStore) GetAll(ctx context.Context, pageToken string, opt ...sessions.SessionStoreOption) (map[string][]byte, string, error) { + bag, err := applyOpts(ctx, opt...) + if err != nil { + return nil, "", err + } + s.mu.RLock() + defer s.mu.RUnlock() + result := make(map[string][]byte) + for k, v := range s.data { + if k.syncID == bag.SyncID { + result[k.key] = v + } + } + return result, "", nil +} + +func applyOpts(ctx context.Context, opt ...sessions.SessionStoreOption) (*sessions.SessionStoreBag, error) { + bag := &sessions.SessionStoreBag{} + for _, o := range opt { + if err := o(ctx, bag); err != nil { + return nil, err + } + } + return bag, nil +} diff --git a/pkg/sync/syncer.go b/pkg/sync/syncer.go index c498b03bf..284862dcb 100644 --- a/pkg/sync/syncer.go +++ b/pkg/sync/syncer.go @@ -475,6 +475,14 @@ func (s *syncer) Sync(ctx context.Context) error { } s.syncID = syncID + // Register the session store keyed by syncID so concurrent syncs each + // resolve to their own backing c1z instead of clobbering a shared pointer. + if s.setSessionStore != nil { + if sessionStore, ok := s.store.(sessions.SessionStore); ok { + s.setSessionStore.SetSessionStore(ctx, syncID, sessionStore) + } + } + // Set the syncID on the wrapper after we have it if syncID == "" { err = ErrNoSyncIDFound @@ -2548,11 +2556,6 @@ func (s *syncer) loadStore(ctx context.Context) error { return err } - // TODO: Remove when pebble supports session store. - sessionStore, ok := store.(sessions.SessionStore) - if s.setSessionStore != nil && ok { - s.setSessionStore.SetSessionStore(ctx, sessionStore) - } s.store = store // Now that s.store is populated, wire the expand progress log's size @@ -2610,6 +2613,10 @@ func (s *syncer) Close(ctx context.Context) error { var errs []error + if s.setSessionStore != nil && s.syncID != "" { + s.setSessionStore.RemoveSessionStore(finalizeCtx, s.syncID) + } + var storeCloseErr error if s.store != nil { storeCloseErr = s.store.Close(finalizeCtx) diff --git a/pkg/types/sessions/sessions.go b/pkg/types/sessions/sessions.go index 841de45fc..daf3db1c0 100644 --- a/pkg/types/sessions/sessions.go +++ b/pkg/types/sessions/sessions.go @@ -73,5 +73,6 @@ func SetSyncIDInContext(ctx context.Context, syncID string) context.Context { } type SetSessionStore interface { - SetSessionStore(ctx context.Context, store SessionStore) + SetSessionStore(ctx context.Context, syncID string, store SessionStore) + RemoveSessionStore(ctx context.Context, syncID string) }