Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 12 additions & 12 deletions cmd/server/api.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
}
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
23 changes: 14 additions & 9 deletions cmd/server/compat_handlers.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package main

import (
"context"
"encoding/json"
"fmt"
"net/http"
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand Down
6 changes: 6 additions & 0 deletions internal/memory/buffer.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
14 changes: 14 additions & 0 deletions internal/memory/buffer_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)}
Expand Down
10 changes: 8 additions & 2 deletions internal/memory/neo4j/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand All @@ -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")
Expand Down
21 changes: 19 additions & 2 deletions internal/memory/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down Expand Up @@ -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})
Expand Down Expand Up @@ -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)
}

Expand Down Expand Up @@ -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) {
Expand Down
2 changes: 1 addition & 1 deletion internal/memory/store.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion sdk/python/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down
Loading