diff --git a/internal/admin/admin.go b/internal/admin/admin.go index 59e2256c..20a96820 100644 --- a/internal/admin/admin.go +++ b/internal/admin/admin.go @@ -1338,7 +1338,11 @@ func (s *Server) handleCapsetPath(w http.ResponseWriter, r *http.Request, capset } } if err := s.Store.AddCapsetMethod(r.Context(), domain.CapsetMethod{CapsetInstanceID: ci.ID, MethodFullName: method.FullName, MCPToolName: toolName, Enabled: true}); err != nil { - writeError(w, http.StatusBadRequest, err.Error()) + statusCode := http.StatusBadRequest + if errors.Is(err, store.ErrMCPToolNameConflict) { + statusCode = http.StatusConflict + } + writeError(w, statusCode, err.Error()) return } } @@ -1440,7 +1444,11 @@ func (s *Server) handleCapsetPath(w http.ResponseWriter, r *http.Request, capset } } if err := s.Store.AddCapsetMethod(r.Context(), domain.CapsetMethod{CapsetInstanceID: ciID, MethodFullName: selected.FullName, MCPToolName: toolName, Enabled: true}); err != nil { - writeError(w, http.StatusBadRequest, err.Error()) + statusCode := http.StatusBadRequest + if errors.Is(err, store.ErrMCPToolNameConflict) { + statusCode = http.StatusConflict + } + writeError(w, statusCode, err.Error()) return } s.logger().Info("capset_method_selected", "capset_id", capsetID, "instance_id", req.InstanceID, "method", selected.FullName, "mcp_tool", toolName) diff --git a/internal/store/mcp_tool_name_test.go b/internal/store/mcp_tool_name_test.go new file mode 100644 index 00000000..a3240e0a --- /dev/null +++ b/internal/store/mcp_tool_name_test.go @@ -0,0 +1,87 @@ +package store + +import ( + "context" + "errors" + "strings" + "testing" + + "octobus/internal/domain" +) + +func TestMigrateReportsExistingMCPToolConflicts(t *testing.T) { + st, err := Open(":memory:") + if err != nil { + t.Fatal(err) + } + defer st.Close() + ctx := context.Background() + if err := st.UpsertService(ctx, domain.Service{ID: "echo", Name: "Echo"}); err != nil { + t.Fatal(err) + } + if err := st.UpsertInstance(ctx, domain.Instance{ID: "echo-instance", ServiceID: "echo", Name: "Echo", Enabled: true}); err != nil { + t.Fatal(err) + } + if err := st.CreateCapset(ctx, domain.Capset{ID: "dev", Name: "Dev", Enabled: true}); err != nil { + t.Fatal(err) + } + if err := st.AddCapsetInstance(ctx, domain.CapsetInstance{ID: "dev:echo-instance", CapsetID: "dev", ServiceID: "echo", InstanceID: "echo-instance", Enabled: true}); err != nil { + t.Fatal(err) + } + if err := st.AddCapsetMethod(ctx, domain.CapsetMethod{CapsetInstanceID: "dev:echo-instance", MethodFullName: "echo.Echo/Call", MCPToolName: "shared_tool", Enabled: true}); err != nil { + t.Fatal(err) + } + if _, err := st.DB().ExecContext(ctx, `DROP INDEX uq_capset_methods_mcp_tool_key`); err != nil { + t.Fatal(err) + } + if _, err := st.DB().ExecContext(ctx, `INSERT INTO capset_methods (id, capset_instance_id, method_full_name, mcp_tool_name, mcp_tool_key, created_at, updated_at) VALUES (?, ?, ?, ?, '', ?, ?)`, "duplicate", "dev:echo-instance", "echo.Echo/Other", "shared_tool", "", ""); err != nil { + t.Fatal(err) + } + if err := st.Migrate(ctx); err == nil || !strings.Contains(err.Error(), "conflicting methods") { + t.Fatalf("migration conflict error = %v", err) + } +} + +func TestAddCapsetMethodEnforcesToolNameUniquenessWithinCapset(t *testing.T) { + st, err := Open(":memory:") + if err != nil { + t.Fatal(err) + } + defer st.Close() + ctx := context.Background() + + if err := st.UpsertService(ctx, domain.Service{ID: "echo", Name: "Echo"}); err != nil { + t.Fatal(err) + } + if err := st.UpsertInstance(ctx, domain.Instance{ID: "echo-instance", ServiceID: "echo", Name: "Echo", Enabled: true}); err != nil { + t.Fatal(err) + } + for _, capsetID := range []string{"dev", "qa"} { + if err := st.CreateCapset(ctx, domain.Capset{ID: capsetID, Name: capsetID, Enabled: true}); err != nil { + t.Fatal(err) + } + if err := st.AddCapsetInstance(ctx, domain.CapsetInstance{ID: capsetID + ":echo-instance", CapsetID: capsetID, ServiceID: "echo", InstanceID: "echo-instance", Enabled: true}); err != nil { + t.Fatal(err) + } + } + if err := st.UpsertInstance(ctx, domain.Instance{ID: "echo-copy", ServiceID: "echo", Name: "Echo Copy", Enabled: true}); err != nil { + t.Fatal(err) + } + if err := st.AddCapsetInstance(ctx, domain.CapsetInstance{ID: "dev:echo-copy", CapsetID: "dev", ServiceID: "echo", InstanceID: "echo-copy", Enabled: true}); err != nil { + t.Fatal(err) + } + + first := domain.CapsetMethod{CapsetInstanceID: "dev:echo-instance", MethodFullName: "echo.Echo/Call", MCPToolName: "shared_tool", Enabled: true} + if err := st.AddCapsetMethod(ctx, first); err != nil { + t.Fatal(err) + } + if err := st.AddCapsetMethod(ctx, domain.CapsetMethod{CapsetInstanceID: "dev:echo-instance", MethodFullName: "echo.Echo/Other", MCPToolName: "shared_tool", Enabled: true}); !errors.Is(err, ErrMCPToolNameConflict) { + t.Fatalf("duplicate tool error = %v", err) + } + if err := st.AddCapsetMethod(ctx, domain.CapsetMethod{CapsetInstanceID: "dev:echo-copy", MethodFullName: "echo.Echo/Call", MCPToolName: "shared_tool", Enabled: true}); !errors.Is(err, ErrMCPToolNameConflict) { + t.Fatalf("cross-instance duplicate tool error = %v", err) + } + if err := st.AddCapsetMethod(ctx, domain.CapsetMethod{CapsetInstanceID: "qa:echo-instance", MethodFullName: "echo.Echo/Call", MCPToolName: "shared_tool", Enabled: true}); err != nil { + t.Fatalf("same tool in another capset error = %v", err) + } +} diff --git a/internal/store/store.go b/internal/store/store.go index 3149b855..58d59c6e 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -22,6 +22,10 @@ type Store struct { db *sql.DB } +var ErrMCPToolNameConflict = errors.New("MCP tool name conflict") + +const mcpToolKeySeparator = "\x1f" + type ServiceInUseError struct { ServiceID string InstanceID string @@ -109,6 +113,27 @@ func (s *Store) Migrate(ctx context.Context) error { if err := addColumnIfMissing(ctx, s.db, "services", "service_root", "TEXT NOT NULL DEFAULT '.'"); err != nil { return err } + if err := addColumnIfMissing(ctx, s.db, "capset_methods", "mcp_tool_key", "TEXT NOT NULL DEFAULT ''"); err != nil { + return err + } + if _, err := s.db.ExecContext(ctx, `UPDATE capset_methods SET mcp_tool_key = ( + SELECT ci.capset_id || char(31) || capset_methods.mcp_tool_name + FROM capset_instances ci + WHERE ci.id = capset_methods.capset_instance_id + ) + WHERE mcp_tool_name <> ''`); err != nil { + return err + } + var duplicateKey string + var duplicateCount int + if err := s.db.QueryRowContext(ctx, `SELECT mcp_tool_key, COUNT(*) FROM capset_methods WHERE mcp_tool_key <> '' GROUP BY mcp_tool_key HAVING COUNT(*) > 1 LIMIT 1`).Scan(&duplicateKey, &duplicateCount); err != nil && !errors.Is(err, sql.ErrNoRows) { + return err + } else if err == nil { + return fmt.Errorf("MCP tool name %q has %d conflicting methods; remove or rename duplicates before restarting", duplicateKey, duplicateCount) + } + if _, err := s.db.ExecContext(ctx, `CREATE UNIQUE INDEX IF NOT EXISTS uq_capset_methods_mcp_tool_key ON capset_methods(mcp_tool_key) WHERE mcp_tool_key <> ''`); err != nil { + return err + } _, err := s.db.ExecContext(ctx, `UPDATE services SET runtime_mode = 'long-running' WHERE runtime_mode = ''`) return err } @@ -782,8 +807,36 @@ func (s *Store) AddCapsetMethod(ctx context.Context, method domain.CapsetMethod) now := time.Now().UTC() method.CreatedAt = now method.UpdatedAt = now - _, err := s.db.ExecContext(ctx, `INSERT INTO capset_methods (id, capset_instance_id, method_full_name, rest_alias, mcp_tool_name, enabled, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, method.ID, method.CapsetInstanceID, method.MethodFullName, method.RestAlias, method.MCPToolName, boolInt(method.Enabled), formatTime(method.CreatedAt), formatTime(method.UpdatedAt)) - return err + + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() + + toolKey := "" + if method.MCPToolName != "" { + var capsetID string + if err := tx.QueryRowContext(ctx, `SELECT capset_id FROM capset_instances WHERE id = ?`, method.CapsetInstanceID).Scan(&capsetID); err != nil { + return err + } + toolKey = capsetID + mcpToolKeySeparator + method.MCPToolName + } + _, err = tx.ExecContext(ctx, `INSERT INTO capset_methods (id, capset_instance_id, method_full_name, rest_alias, mcp_tool_name, mcp_tool_key, enabled, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, method.ID, method.CapsetInstanceID, method.MethodFullName, method.RestAlias, method.MCPToolName, toolKey, boolInt(method.Enabled), formatTime(method.CreatedAt), formatTime(method.UpdatedAt)) + if err != nil { + var existingID string + if idErr := tx.QueryRowContext(ctx, `SELECT id FROM capset_methods WHERE id = ?`, method.ID).Scan(&existingID); idErr == nil { + return err + } + if toolKey != "" { + var existingToolID string + if keyErr := tx.QueryRowContext(ctx, `SELECT id FROM capset_methods WHERE mcp_tool_key = ?`, toolKey).Scan(&existingToolID); keyErr == nil { + return ErrMCPToolNameConflict + } + } + return err + } + return tx.Commit() } func (s *Store) GetCapsetMethod(ctx context.Context, capsetInstanceID, methodFullName string) (domain.CapsetMethod, error) { @@ -1257,6 +1310,7 @@ CREATE TABLE IF NOT EXISTS capset_methods ( method_full_name TEXT NOT NULL, rest_alias TEXT NOT NULL DEFAULT '', mcp_tool_name TEXT NOT NULL DEFAULT '', + mcp_tool_key TEXT NOT NULL DEFAULT '', enabled INTEGER NOT NULL DEFAULT 1, created_at TEXT NOT NULL, updated_at TEXT NOT NULL, diff --git a/internal/store/store_test.go b/internal/store/store_test.go index 9edc83c9..bdbfd1db 100644 --- a/internal/store/store_test.go +++ b/internal/store/store_test.go @@ -935,14 +935,14 @@ func TestFindToolAmbiguousAndStreamingCustomNames(t *testing.T) { if err := s.AddCapsetInstance(ctx, domain.CapsetInstance{ID: "dev:echo-copy", CapsetID: "dev", ServiceID: "echo", InstanceID: "echo-copy", IncludeAllMethods: false, Enabled: true}); err != nil { t.Fatal(err) } - if err := s.AddCapsetMethod(ctx, domain.CapsetMethod{CapsetInstanceID: "dev:echo-copy", MethodFullName: "echo.v1.EchoService/Echo", MCPToolName: "echo_tool", Enabled: true}); err != nil { - t.Fatal(err) + if err := s.AddCapsetMethod(ctx, domain.CapsetMethod{CapsetInstanceID: "dev:echo-copy", MethodFullName: "echo.v1.EchoService/Echo", MCPToolName: "echo_tool", Enabled: true}); !errors.Is(err, ErrMCPToolNameConflict) { + t.Fatalf("expected duplicate MCP tool error, got %v", err) } - if _, err := s.FindTool(ctx, "dev", "echo_tool"); err == nil || !strings.Contains(err.Error(), "ambiguous MCP tool name") { - t.Fatalf("expected ambiguous MCP tool error, got %v", err) + if _, err := s.FindTool(ctx, "dev", "echo_tool"); err != nil { + t.Fatalf("FindTool existing unique name error = %v", err) } - if exists, err := s.MCPToolNameExists(ctx, "dev", "echo_tool"); err == nil || !strings.Contains(err.Error(), "ambiguous MCP tool name") || exists { - t.Fatalf("MCPToolNameExists ambiguous exists=%v err=%v", exists, err) + if exists, err := s.MCPToolNameExists(ctx, "dev", "echo_tool"); err != nil || !exists { + t.Fatalf("MCPToolNameExists exists=%v err=%v", exists, err) } if err := s.AddCapsetMethod(ctx, domain.CapsetMethod{CapsetInstanceID: "dev:echo-copy", MethodFullName: "echo.v1.EchoService/ServerStream", MCPToolName: "stream_custom", Enabled: true}); err != nil {