From 76bd7560bfe6c9fb32c124637c2f829b04d009e6 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Mon, 20 Jul 2026 14:27:15 +0000 Subject: [PATCH 1/2] fix: harden memory buffer, reranker init, and tenant-scoped bulk delete - Guard MessageBuffer flush when Neo4j backend is unavailable (typed-nil fix) - Clear reranker interface on init failure to prevent search panics - Scope BulkDeleteByFilter by auth-bound tenant_id in Neo4j and service layer - Route memory mutation/compat handlers through requestContextWithTenant Co-authored-by: Himan --- cmd/server/api.go | 24 ++++++++++++------------ cmd/server/compat_handlers.go | 23 ++++++++++++++--------- internal/memory/buffer.go | 6 ++++++ internal/memory/buffer_test.go | 14 ++++++++++++++ internal/memory/neo4j/client.go | 10 ++++++++-- internal/memory/service.go | 21 +++++++++++++++++++-- internal/memory/store.go | 2 +- 7 files changed, 74 insertions(+), 26 deletions(-) diff --git a/cmd/server/api.go b/cmd/server/api.go index b671e5f4..04312751 100644 --- a/cmd/server/api.go +++ b/cmd/server/api.go @@ -2747,7 +2747,7 @@ func (s *APIServer) processMemoryHandler(w http.ResponseWriter, r *http.Request) CustomInstructions: req.CustomInstructions, } - created, err := s.memSvc.CreateMemoryWithOptions(context.Background(), mem, req.SkipProcessing) + created, err := s.memSvc.CreateMemoryWithOptions(requestContextWithTenant(r), mem, req.SkipProcessing) if err != nil { http.Error(w, "Failed to process memory", http.StatusInternalServerError) return @@ -2902,12 +2902,12 @@ func (s *APIServer) updateMemoryHandler(w http.ResponseWriter, r *http.Request) return } - if err := s.memSvc.UpdateMemory(context.Background(), memoryID, req.Content, req.Metadata); err != nil { + if err := s.memSvc.UpdateMemory(requestContextWithTenant(r), memoryID, req.Content, req.Metadata); err != nil { safeHTTPError(w, r, err, http.StatusInternalServerError) return } - mem, _ := s.memSvc.GetMemory(context.Background(), memoryID) + mem, _ := s.memSvc.GetMemory(requestContextWithTenant(r), memoryID) s.emitSSE(getTenantID(r), "memory.updated", mem) json.NewEncoder(w).Encode(mem) } @@ -2929,7 +2929,7 @@ func (s *APIServer) getMemoryHistoryHandler(w http.ResponseWriter, r *http.Reque vars := mux.Vars(r) memoryID := vars["memoryID"] - history, err := s.memSvc.GetMemoryHistory(context.Background(), memoryID) + history, err := s.memSvc.GetMemoryHistory(requestContextWithTenant(r), memoryID) if err != nil { safeHTTPError(w, r, err, http.StatusInternalServerError) return @@ -2993,14 +2993,14 @@ func (s *APIServer) batchCreateMemoriesHandler(w http.ResponseWriter, r *http.Re return } - tenantID := getTenantID(r) + tenantID := effectiveTenantID(r) for _, mem := range req.Memories { if tenantID != "" { mem.TenantID = tenantID } } - created, err := s.memSvc.BatchCreateMemories(context.Background(), req.Memories) + created, err := s.memSvc.BatchCreateMemories(requestContextWithTenant(r), req.Memories) if err != nil { safeHTTPError(w, r, err, http.StatusInternalServerError) return @@ -3033,7 +3033,7 @@ func (s *APIServer) batchUpdateMemoriesHandler(w http.ResponseWriter, r *http.Re return } - if err := s.memSvc.BatchUpdateMemories(context.Background(), &req); err != nil { + if err := s.memSvc.BatchUpdateMemories(requestContextWithTenant(r), &req); err != nil { safeHTTPError(w, r, err, http.StatusInternalServerError) return } @@ -3079,7 +3079,7 @@ func (s *APIServer) bulkDeleteHandler(w http.ResponseWriter, r *http.Request) { return } - count, err := s.memSvc.BulkDeleteByFilter(context.Background(), &req) + count, err := s.memSvc.BulkDeleteByFilter(requestContextWithTenant(r), &req) if err != nil { safeHTTPError(w, r, err, http.StatusInternalServerError) return @@ -3113,7 +3113,7 @@ func (s *APIServer) resetMemoriesHandler(w http.ResponseWriter, r *http.Request) Category: category, } - count, err := s.memSvc.BulkDeleteByFilter(context.Background(), req) + count, err := s.memSvc.BulkDeleteByFilter(requestContextWithTenant(r), req) if err != nil { safeHTTPError(w, r, err, http.StatusInternalServerError) return @@ -3191,7 +3191,7 @@ func (s *APIServer) createMemoryFeedbackHandler(w http.ResponseWriter, r *http.R func (s *APIServer) listFeedbackHandler(w http.ResponseWriter, r *http.Request) { memID := r.URL.Query().Get("memory_id") if memID != "" { - history, _ := s.memSvc.GetMemoryHistory(context.Background(), memID) + history, _ := s.memSvc.GetMemoryHistory(requestContextWithTenant(r), memID) var feedback []types.MemoryHistory for _, h := range history { if h.Action == types.HistoryActionFeedback { @@ -5558,7 +5558,7 @@ func (s *APIServer) getMemoryVersionsHandler(w http.ResponseWriter, r *http.Requ vars := mux.Vars(r) memoryID := vars["memoryID"] - history, err := s.memSvc.GetMemoryHistory(context.Background(), memoryID) + history, err := s.memSvc.GetMemoryHistory(requestContextWithTenant(r), memoryID) if err != nil { safeHTTPError(w, r, err, http.StatusInternalServerError) return @@ -5585,7 +5585,7 @@ func (s *APIServer) restoreMemoryVersionHandler(w http.ResponseWriter, r *http.R return } - history, err := s.memSvc.GetMemoryHistory(context.Background(), memoryID) + history, err := s.memSvc.GetMemoryHistory(requestContextWithTenant(r), memoryID) if err != nil { safeHTTPError(w, r, err, http.StatusInternalServerError) return diff --git a/cmd/server/compat_handlers.go b/cmd/server/compat_handlers.go index c4137a3d..cb207143 100644 --- a/cmd/server/compat_handlers.go +++ b/cmd/server/compat_handlers.go @@ -1,7 +1,6 @@ package main import ( - "context" "encoding/json" "fmt" "net/http" @@ -131,10 +130,10 @@ func (s *APIServer) v3AddMemoriesHandler(w http.ResponseWriter, r *http.Request) CreatedAt: time.Now(), UpdatedAt: time.Now(), } - if tenantID := getTenantID(r); tenantID != "" { + if tenantID := effectiveTenantID(r); tenantID != "" { mem.TenantID = tenantID } - created, err := s.memSvc.CreateMemoryWithOptions(context.Background(), mem, req.SkipProcessing) + created, err := s.memSvc.CreateMemoryWithOptions(requestContextWithTenant(r), mem, req.SkipProcessing) if err != nil { safeHTTPError(w, r, fmt.Errorf("v3 add memory: %w", err), http.StatusInternalServerError) return @@ -204,8 +203,9 @@ func (s *APIServer) v3SearchMemoriesHandler(w http.ResponseWriter, r *http.Reque Category: firstCategory(req.Categories), Mode: "hybrid", Rerank: req.Rerank, + TenantID: effectiveTenantID(r), } - results, err := s.memSvc.SearchMemories(context.Background(), searchReq) + results, err := s.memSvc.SearchMemories(requestContextWithTenant(r), searchReq) if err != nil { safeHTTPError(w, r, fmt.Errorf("v3 search memories: %w", err), http.StatusInternalServerError) return @@ -243,7 +243,7 @@ func (s *APIServer) v3ListMemoriesHandler(w http.ResponseWriter, r *http.Request req.PageSize = 100 } - memories, err := s.listScopedMemories(req.UserID, req.OrgID) + memories, err := s.listScopedMemories(r, req.UserID, req.OrgID) if err != nil { safeHTTPError(w, r, fmt.Errorf("v3 list memories: %w", err), http.StatusInternalServerError) return @@ -331,14 +331,19 @@ func (s *APIServer) createImportHandler(w http.ResponseWriter, r *http.Request) }) } -func (s *APIServer) listScopedMemories(userID, orgID string) ([]*types.Memory, error) { +func (s *APIServer) listScopedMemories(r *http.Request, userID, orgID string) ([]*types.Memory, error) { + ctx := requestContextWithTenant(r) + tenantID := effectiveTenantID(r) + if tenantID != "" { + return s.memSvc.GetMemoriesByTenant(ctx, tenantID, 1000) + } if userID != "" { - return s.memSvc.GetMemoriesByUser(context.Background(), userID) + return s.memSvc.GetMemoriesByUser(ctx, userID) } if orgID != "" { - return s.memSvc.GetMemoriesByOrg(context.Background(), orgID) + return s.memSvc.GetMemoriesByOrg(ctx, orgID) } - return s.memSvc.GetAllMemories(context.Background()) + return s.memSvc.GetMemoriesByTenant(ctx, "default", 1000) } func v3Content(req v3MemoryInput) string { diff --git a/internal/memory/buffer.go b/internal/memory/buffer.go index d23d84f1..e4ea99cc 100644 --- a/internal/memory/buffer.go +++ b/internal/memory/buffer.go @@ -79,6 +79,12 @@ func (mb *MessageBuffer) flushSession(sessionID string) error { return nil } + if mb.neo4j == nil { + mb.messages[sessionID] = nil + delete(mb.messages, sessionID) + return nil + } + for _, msg := range msgs { if err := mb.neo4j.AddMessage(sessionID, msg); err != nil { fmt.Printf("warn: buffer flush message %s: %v\n", msg.ID, err) diff --git a/internal/memory/buffer_test.go b/internal/memory/buffer_test.go index c7999487..e37ab6ba 100644 --- a/internal/memory/buffer_test.go +++ b/internal/memory/buffer_test.go @@ -39,6 +39,20 @@ func TestMessageBuffer_Add(t *testing.T) { t.Errorf("expected buffer length 1, got %d", buf.Len()) } } +func TestMessageBuffer_NilBackend(t *testing.T) { + buf := NewMessageBuffer(2, time.Hour, nil) + + if err := buf.Add(types.Message{ID: "1", SessionID: "s1"}); err != nil { + t.Fatalf("Add failed: %v", err) + } + if err := buf.Add(types.Message{ID: "2", SessionID: "s1"}); err != nil { + t.Fatalf("Add failed: %v", err) + } + + if buf.Len() != 0 { + t.Fatalf("expected flush to drop messages when backend nil, got length %d", buf.Len()) + } +} func TestMessageBuffer_FlushOnSize(t *testing.T) { mock := &mockNeo4j{messages: make(map[string][]types.Message)} diff --git a/internal/memory/neo4j/client.go b/internal/memory/neo4j/client.go index fee4de17..ff4f6a62 100644 --- a/internal/memory/neo4j/client.go +++ b/internal/memory/neo4j/client.go @@ -1754,7 +1754,10 @@ func (c *Client) GetExpiredMemories() ([]*types.Memory, error) { return memories, nil } -func (c *Client) BulkDeleteByFilter(userID, orgID, category, agentID string) (int, error) { +func (c *Client) BulkDeleteByFilter(tenantID, userID, orgID, category, agentID string) (int, error) { + if tenantID == "" { + return 0, fmt.Errorf("tenant_id required") + } ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second) defer cancel() @@ -1764,7 +1767,10 @@ func (c *Client) BulkDeleteByFilter(userID, orgID, category, agentID string) (in defer session.Close(ctx) var conditions []string - params := map[string]interface{}{} + params := map[string]interface{}{ + "tenant_id": tenantID, + } + conditions = append(conditions, "m.tenant_id = $tenant_id") if userID != "" { conditions = append(conditions, "m.user_id = $user_id") diff --git a/internal/memory/service.go b/internal/memory/service.go index 938f4676..052712f8 100644 --- a/internal/memory/service.go +++ b/internal/memory/service.go @@ -185,7 +185,13 @@ func NewService(cfg *config.Config) (*Service, error) { } else if vec != nil { svc.vector = vec } - svc.msgBuffer = NewMessageBuffer(cfg.App.MessageBuffer, cfg.App.BufferTimeout, neo) + var msgNeo interface { + AddMessage(sessionID string, msg types.Message) error + } + if neo != nil { + msgNeo = neo + } + svc.msgBuffer = NewMessageBuffer(cfg.App.MessageBuffer, cfg.App.BufferTimeout, msgNeo) if cfg.LLM.APIKey != "" { llmCfg := &llm.Config{ Provider: llm.ProviderType(cfg.LLM.Provider), @@ -250,6 +256,7 @@ func NewService(cfg *config.Config) (*Service, error) { svc.reranker, rerankerErr = reranker.NewProvider(cfg.Reranker, svc.llmClient) if rerankerErr != nil { log.Printf("warning: reranker unavailable: %v", rerankerErr) + svc.reranker = nil } svc.compStats = &CompressionStats{} svc.privacyFilter = privacy.NewFilter(privacy.FilterConfig{Enabled: cfg.Privacy.Enabled}) @@ -1773,6 +1780,9 @@ func (s *Service) GetMemoryHistory(ctx context.Context, id string) ([]types.Memo if s.graph == nil { return nil, nil } + if _, err := s.GetMemory(ctx, id); err != nil { + return nil, err + } return s.graph.GetMemoryHistory(id) } @@ -1893,7 +1903,14 @@ func (s *Service) BulkDeleteByFilter(ctx context.Context, req *types.BatchDelete if req == nil { return 0, nil } - return s.graph.BulkDeleteByFilter(req.UserID, req.OrgID, req.Category, req.AgentID) + tid := tenant.IDFromContext(ctx) + if tid == "" { + tid = s.defaultTenantID + } + if tid == "" { + return 0, fmt.Errorf("service: bulk delete requires tenant context") + } + return s.graph.BulkDeleteByFilter(tid, req.UserID, req.OrgID, req.Category, req.AgentID) } func (s *Service) AddFeedback(ctx context.Context, fb *types.Feedback) (*types.Feedback, error) { diff --git a/internal/memory/store.go b/internal/memory/store.go index 3dab415d..510263e8 100644 --- a/internal/memory/store.go +++ b/internal/memory/store.go @@ -28,7 +28,7 @@ type GraphStore interface { GetMemoryHistory(memID string) ([]types.MemoryHistory, error) AdvancedSearch(filters *types.SearchFilters) ([]*types.Memory, error) - BulkDeleteByFilter(userID, orgID, category, agentID string) (int, error) + BulkDeleteByFilter(tenantID, userID, orgID, category, agentID string) (int, error) CreateSession(agentID string, metadata map[string]interface{}) (*types.Session, error) ListSessions() ([]*types.Session, error) From fbb37d69067ec19e83b3e4db4ecc183de9babe01 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Mon, 27 Jul 2026 14:18:41 +0000 Subject: [PATCH 2/2] fix(ci): exclude Python SDK README from ruff format check README.md contains markdown code examples that ruff format would rewrite, causing false CI failures on master push runs. Co-authored-by: Himan --- sdk/python/pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sdk/python/pyproject.toml b/sdk/python/pyproject.toml index a5f02f92..8461b2b4 100644 --- a/sdk/python/pyproject.toml +++ b/sdk/python/pyproject.toml @@ -55,7 +55,7 @@ exclude = ["tests*", "examples*"] [tool.ruff] target-version = "py39" -exclude = ["build", "dist", "*.egg-info"] +exclude = ["build", "dist", "*.egg-info", "README.md"] [tool.ruff.lint] select = ["E", "F", "I"]