diff --git a/internal/mcp/mcp.go b/internal/mcp/mcp.go index e1fb4d16..36540ac4 100644 --- a/internal/mcp/mcp.go +++ b/internal/mcp/mcp.go @@ -98,6 +98,7 @@ func ensureImplicitSessionWithCWD(s *store.Store, sessionID, project string) err var ProfileAgent = map[string]bool{ "mem_save": true, // proactive save — referenced 17 times across protocols "mem_search": true, // search past memories — referenced 6 times + "mem_find_project": true, // find projects by memory content "mem_context": true, // recent context from previous sessions — referenced 10 times "mem_session_summary": true, // end-of-session summary — referenced 16 times "mem_session_start": true, // register session start @@ -295,6 +296,28 @@ func registerTools(srv *server.MCPServer, s *store.Store, cfg MCPConfig, allowli ) } + // ─── mem_find_project ───────────────────────────────────────────── + if shouldRegister("mem_find_project", allowlist) { + srv.AddTool( + mcp.NewTool("mem_find_project", + mcp.WithDescription("Search for projects containing relevant memories. Use this when you don't know which project holds a past decision. It returns the top matching projects, their match counts, and rank. You can then use mem_search with a specific project name to read those memories."), + mcp.WithTitleAnnotation("Find Projects"), + mcp.WithReadOnlyHintAnnotation(true), + mcp.WithDestructiveHintAnnotation(false), + mcp.WithIdempotentHintAnnotation(true), + mcp.WithOpenWorldHintAnnotation(false), + mcp.WithString("query", + mcp.Required(), + mcp.Description("Search query — natural language or keywords to find across all projects"), + ), + mcp.WithString("match_mode", + mcp.Description("Token matching: \"all\" (default — every token must match, FTS5 AND) or \"any\" (any token matches)."), + ), + ), + handleFindProject(s, cfg), + ) + } + // ─── mem_save (profile: agent, core — always in context) ─────────── if shouldRegister("mem_save", allowlist) { srv.AddTool( @@ -1148,6 +1171,39 @@ func handleSearch(s *store.Store, cfg MCPConfig, activity *SessionActivity) serv } } +func handleFindProject(s *store.Store, cfg MCPConfig) server.ToolHandlerFunc { + return func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { + query, _ := req.GetArguments()["query"].(string) + matchMode, _ := req.GetArguments()["match_mode"].(string) + + if query == "" { + return mcp.NewToolResultError("query is required"), nil + } + if matchMode != "" && matchMode != "all" && matchMode != "any" { + return mcp.NewToolResultError(fmt.Sprintf("invalid match_mode %q: must be \"all\" or \"any\"", matchMode)), nil + } + + limit := 10 // Fix limit as requested by minimalist approach + matches, err := s.SearchProjects(query, matchMode, limit) + if err != nil { + return mcp.NewToolResultError("Project search failed: " + err.Error()), nil + } + + if len(matches) == 0 { + return mcp.NewToolResultText(fmt.Sprintf("No projects found matching %q.", query)), nil + } + + var b strings.Builder + fmt.Fprintf(&b, "Found %d project(s) matching %q:\n", len(matches), query) + for _, m := range matches { + fmt.Fprintf(&b, "- %s (%d matches, rank: %.2f)\n", m.Project, m.MatchCount, m.TopRank) + } + b.WriteString("\nUse mem_search with project: \"\" to explore these memories.") + + return mcp.NewToolResultText(b.String()), nil + } +} + func handlePin(s *store.Store, pinned bool) server.ToolHandlerFunc { return func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { id := int64(intArg(req, "id", 0)) diff --git a/internal/mcp/mcp_test.go b/internal/mcp/mcp_test.go index 89fc21ed..9e6b643b 100644 --- a/internal/mcp/mcp_test.go +++ b/internal/mcp/mcp_test.go @@ -1611,7 +1611,7 @@ func TestResolveToolsAgentProfile(t *testing.T) { } expectedTools := []string{ - "mem_save", "mem_search", "mem_context", "mem_session_summary", + "mem_save", "mem_search", "mem_find_project", "mem_context", "mem_session_summary", "mem_session_start", "mem_session_end", "mem_get_observation", "mem_suggest_topic_key", "mem_capture_passive", "mem_save_prompt", "mem_update", // skills explicitly say "use mem_update when you have an exact ID to correct" @@ -2254,7 +2254,7 @@ func TestNewServerWithToolsNilRegistersAll(t *testing.T) { tools := srv.ListTools() allTools := []string{ - "mem_save", "mem_search", "mem_context", "mem_session_summary", + "mem_save", "mem_search", "mem_find_project", "mem_context", "mem_session_summary", "mem_session_start", "mem_session_end", "mem_get_observation", "mem_suggest_topic_key", "mem_capture_passive", "mem_save_prompt", "mem_update", "mem_delete", "mem_stats", "mem_timeline", "mem_merge_projects", @@ -2364,14 +2364,14 @@ func TestNewServerBackwardsCompatible(t *testing.T) { srv := NewServer(s) tools := srv.ListTools() - // 18 agent + 4 admin = 22 total. - if len(tools) != 22 { - t.Errorf("NewServer should register all 22 tools, got %d", len(tools)) + // 19 agent + 4 admin = 23 total. + if len(tools) != 23 { + t.Errorf("NewServer should register all 23 tools, got %d", len(tools)) } } func TestProfileConsistency(t *testing.T) { - // Verify that agent + admin = all 22 tools + // Verify that agent + admin = all 23 tools combined := make(map[string]bool) for tool := range ProfileAgent { combined[tool] = true @@ -2380,9 +2380,9 @@ func TestProfileConsistency(t *testing.T) { combined[tool] = true } - // 18 agent + 4 admin = 22 total. - if len(combined) != 22 { - t.Errorf("agent + admin should cover all 22 tools, got %d", len(combined)) + // 19 agent + 4 admin = 23 total. + if len(combined) != 23 { + t.Errorf("agent + admin should cover all 23 tools, got %d", len(combined)) } // Verify no overlap between profiles @@ -2710,9 +2710,9 @@ func TestNewServerWithConfig(t *testing.T) { t.Fatal("expected MCP server instance") } tools := srv.ListTools() - // Should have all 22 tools (18 agent + 4 admin). - if len(tools) != 22 { - t.Errorf("NewServerWithConfig should register all 22 tools, got %d", len(tools)) + // Should have all 23 tools (19 agent + 4 admin). + if len(tools) != 23 { + t.Errorf("NewServerWithConfig should register all 23 tools, got %d", len(tools)) } } diff --git a/internal/store/store.go b/internal/store/store.go index 9c6537b9..83049f94 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -182,6 +182,12 @@ type SearchOptions struct { MatchMode string `json:"match_mode,omitempty"` // "all" (default) | "any" } +type ProjectMatch struct { + Project string + MatchCount int + TopRank float64 +} + type AddObservationParams struct { SessionID string `json:"session_id"` Type string `json:"type"` @@ -3235,6 +3241,59 @@ func (s *Store) Search(query string, opts SearchOptions) ([]SearchResult, error) return results, nil } +// ─── Search Projects ──────────────────────────────────────────────────────── + +// SearchProjects groups FTS5 search results by project to help route ambiguous searches. +func (s *Store) SearchProjects(query string, matchMode string, limit int) ([]ProjectMatch, error) { + if limit <= 0 { + limit = 10 + } + if limit > 50 { + limit = 50 + } + + var ftsQuery string + if matchMode == "any" { + ftsQuery = sanitizeFTSCandidates(query) + } else { + ftsQuery = sanitizeFTS(query) + } + if ftsQuery == "" { + return []ProjectMatch{}, nil + } + + sqlQ := ` + SELECT project, COUNT(id) as match_count, MIN(rank) as top_rank + FROM ( + SELECT o.project, o.id, observations_fts.rank as rank + FROM observations_fts + JOIN observations o ON o.id = observations_fts.rowid + WHERE observations_fts MATCH ? AND o.deleted_at IS NULL AND o.project != '' + ) + GROUP BY project + ORDER BY top_rank ASC, match_count DESC, project ASC + LIMIT ? + ` + rows, err := s.queryItHook(s.db, sqlQ, ftsQuery, limit) + if err != nil { + return nil, fmt.Errorf("search projects: %w", err) + } + defer rows.Close() + + var matches []ProjectMatch + for rows.Next() { + var p ProjectMatch + if err := rows.Scan(&p.Project, &p.MatchCount, &p.TopRank); err != nil { + return nil, err + } + matches = append(matches, p) + } + if err := rows.Err(); err != nil { + return nil, err + } + return matches, nil +} + // ─── Stats ─────────────────────────────────────────────────────────────────── func (s *Store) Stats() (*Stats, error) { diff --git a/internal/store/store_test.go b/internal/store/store_test.go index 5ed55ca6..5a7eca76 100644 --- a/internal/store/store_test.go +++ b/internal/store/store_test.go @@ -8830,3 +8830,77 @@ func TestSanitizeFTS(t *testing.T) { }) } } + +func TestSearchProjects(t *testing.T) { + st := newTestStore(t) + + st.CreateSession("s1", "project-a", "/tmp/a") + st.CreateSession("s2", "project-b", "/tmp/b") + st.CreateSession("s3", "project-c", "/tmp/c") + + // Seed 3 observations for project A + _, err := st.AddObservation(AddObservationParams{ + SessionID: "s1", Type: "bugfix", Title: "Fix auth token expiration", + Content: "The auth middleware was dropping tokens", Project: "project-a", + }) + if err != nil { t.Fatal(err) } + _, err = st.AddObservation(AddObservationParams{ + SessionID: "s1", Type: "bugfix", Title: "Auth token validation", + Content: "Middleware should validate auth tokens", Project: "project-a", + }) + if err != nil { t.Fatal(err) } + _, err = st.AddObservation(AddObservationParams{ + SessionID: "s1", Type: "bugfix", Title: "Minor fix", + Content: "Just a minor auth fix in the middleware", Project: "project-a", + }) + if err != nil { t.Fatal(err) } + + // Seed 1 highly relevant observation for project B + _, err = st.AddObservation(AddObservationParams{ + SessionID: "s2", Type: "bugfix", Title: "Auth middleware completely rewritten", + Content: "Auth middleware auth middleware auth middleware tokens", Project: "project-b", + }) + if err != nil { t.Fatal(err) } + + // Seed an irrelevant observation for project C + _, err = st.AddObservation(AddObservationParams{ + SessionID: "s3", Type: "feature", Title: "Database migration", + Content: "Added new tables", Project: "project-c", + }) + if err != nil { t.Fatal(err) } + + // Force FTS sync if async (test setup normally does this, but just in case) + // We'll just search directly. + + matches, err := st.SearchProjects("auth middleware", "all", 10) + if err != nil { + t.Fatalf("SearchProjects failed: %v", err) + } + + if len(matches) != 2 { + t.Fatalf("Expected 2 projects, got %d: %+v", len(matches), matches) + } + + // project-a should have 3 matches. project-b should have 1 match. + // project-b has more occurrences of the terms, so its top_rank might be better (more negative). + // Let's assert on the project names and counts. + hasA := false + hasB := false + for _, m := range matches { + if m.Project == "project-a" { + hasA = true + if m.MatchCount != 3 { + t.Errorf("project-a: expected 3 matches, got %d", m.MatchCount) + } + } + if m.Project == "project-b" { + hasB = true + if m.MatchCount != 1 { + t.Errorf("project-b: expected 1 match, got %d", m.MatchCount) + } + } + } + if !hasA || !hasB { + t.Errorf("Missing expected projects in results: %+v", matches) + } +}