Skip to content
Open
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
63 changes: 63 additions & 0 deletions internal/mcp/mcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -295,6 +296,31 @@ 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)."),
),
mcp.WithString("scope",
mcp.Description("Filter search results by scope: \"project\" (only team/project workspace memories), \"personal\" (personal logs/diary), or \"all\" (default — search across both)."),
),
),
handleFindProject(s, cfg),
)
}

// ─── mem_save (profile: agent, core — always in context) ───────────
if shouldRegister("mem_save", allowlist) {
srv.AddTool(
Expand Down Expand Up @@ -1148,6 +1174,43 @@ 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)
scope, _ := req.GetArguments()["scope"].(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
}
if scope != "" && scope != "all" && scope != "project" && scope != "personal" {
return mcp.NewToolResultError(fmt.Sprintf("invalid scope %q: must be \"all\", \"project\", or \"personal\"", scope)), nil
}

limit := 10 // Fix limit as requested by minimalist approach
matches, err := s.SearchProjects(query, matchMode, scope, 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: %g)\n", m.Project, m.MatchCount, m.TopRank)
}
b.WriteString("\nUse mem_search with project: \"<name>\" 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))
Expand Down
167 changes: 155 additions & 12 deletions internal/mcp/mcp_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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",
Comment thread
coderabbitai[bot] marked this conversation as resolved.
"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",
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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))
}
}

Expand Down Expand Up @@ -7452,3 +7452,146 @@ func TestHandleSearch_MatchModeInvalidError(t *testing.T) {
t.Fatalf("parameter-validation error must not contain query-advice suffix \"Try simpler keywords\", got: %s", text)
}
}

func TestHandleFindProject(t *testing.T) {
s := newMCPTestStore(t)

s.CreateSession("s1", "project-one", "/tmp/one")
s.CreateSession("s2", "project-two", "/tmp/two")

// Insert some observations in different projects to test SearchProjects
_, err := s.AddObservation(store.AddObservationParams{
SessionID: "s1", Type: "bugfix", Title: "search test one",
Content: "project one search test", Project: "project-one",
})
if err != nil { t.Fatal(err) }
_, err = s.AddObservation(store.AddObservationParams{
SessionID: "s2", Type: "bugfix", Title: "search test two",
Content: "project two search test", Project: "project-two",
})
if err != nil { t.Fatal(err) }

// Insert a personal observation in project-one to test scope filtering
_, err = s.AddObservation(store.AddObservationParams{
SessionID: "s1", Type: "bugfix", Title: "personal search test",
Content: "project one personal thoughts", Project: "project-one", Scope: "personal",
})
if err != nil { t.Fatal(err) }

h := handleFindProject(s, MCPConfig{})

tests := []struct {
name string
query string
matchMode string
scope string
expectError bool
errorContains string
expectText string
}{
{
name: "success exact match",
query: "project one",
matchMode: "",
expectText: "Found 1 project(s)",
},
{
name: "success match any",
query: "project one test",
matchMode: "any",
expectText: "project-two", // Both contain test
},
{
name: "missing query",
query: "",
matchMode: "",
expectError: true,
errorContains: "query is required",
},
{
name: "invalid match_mode",
query: "test",
matchMode: "invalid-mode",
expectError: true,
errorContains: "invalid match_mode",
},
{
name: "invalid scope",
query: "test",
matchMode: "",
scope: "invalid-scope",
expectError: true,
errorContains: "invalid scope",
},
{
name: "no results",
query: "nonexistentstringthatwillnevermatch",
matchMode: "",
expectText: "No projects found matching",
},
{
name: "scope project filters out personal ones",
query: "thoughts",
matchMode: "",
scope: "project",
expectText: "No projects found matching",
},
{
name: "scope personal finds only personal ones",
query: "thoughts",
matchMode: "",
scope: "personal",
expectText: "project-one",
},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
args := map[string]any{"query": tc.query}
if tc.matchMode != "" {
args["match_mode"] = tc.matchMode
}
if tc.scope != "" {
args["scope"] = tc.scope
}
req := mcppkg.CallToolRequest{Params: mcppkg.CallToolParams{Arguments: args}}
res, err := h(context.Background(), req)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}

if tc.expectError {
if !res.IsError {
t.Fatalf("expected error, got success")
}
text := callResultText(t, res)
if !strings.Contains(text, tc.errorContains) {
t.Fatalf("expected error containing %q, got %q", tc.errorContains, text)
}
} else {
if res.IsError {
t.Fatalf("expected success, got error: %s", callResultText(t, res))
}
text := callResultText(t, res)
if tc.expectText != "" && !strings.Contains(text, tc.expectText) {
t.Fatalf("expected text containing %q, got %q", tc.expectText, text)
}
}
})
}

// Test error from store (e.g., closed db)
s.Close()
req := mcppkg.CallToolRequest{Params: mcppkg.CallToolParams{Arguments: map[string]any{"query": "test"}}}
res, err := h(context.Background(), req)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !res.IsError {
t.Fatalf("expected tool error due to closed store")
}
text := callResultText(t, res)
if !strings.Contains(text, "Project search failed") {
t.Fatalf("expected error text for store failure, got: %s", text)
}
}
71 changes: 71 additions & 0 deletions internal/store/store.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"`
Expand Down Expand Up @@ -3235,6 +3241,71 @@ 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, scope 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
}

var args []any
args = append(args, ftsQuery)

scopeFilter := ""
if scope != "" && scope != "all" {
scopeFilter = " AND o.scope = ?"
args = append(args, normalizeScope(scope))
}

args = append(args, limit)

sqlQ := fmt.Sprintf(`
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 != ''%s
)
GROUP BY project
ORDER BY top_rank ASC, match_count DESC, project ASC
LIMIT ?
`, scopeFilter)

rows, err := s.queryItHook(s.db, sqlQ, args...)
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) {
Expand Down
Loading
Loading