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
17 changes: 11 additions & 6 deletions internal/api/health.go
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,10 @@ type ServerHealthResponse struct {
Body ServerHealth
}

type healthRoutes struct {
monitor contracts.MCPHealthMonitor
}

// ToAPIType can be used to convert a wrapped domain type to an API-safe type.
func (d DomainServerHealth) ToAPIType() (ServerHealth, error) {
status, err := parseHealthStatus(d.Status)
Expand All @@ -83,6 +87,7 @@ func (d DomainServerHealth) ToAPIType() (ServerHealth, error) {
// RegisterHealthRoutes sets up health-related API endpoint routes.
func RegisterHealthRoutes(routerAPI huma.API, monitor contracts.MCPHealthMonitor, apiPathPrefix string) {
healthAPI := huma.NewGroup(routerAPI, apiPathPrefix)
routes := &healthRoutes{monitor: monitor}
tags := []string{"Health"}

huma.Register(
Expand All @@ -95,7 +100,7 @@ func RegisterHealthRoutes(routerAPI huma.API, monitor contracts.MCPHealthMonitor
Tags: tags,
},
func(ctx context.Context, _ *struct{}) (*ServersHealthResponse, error) {
return handleHealthServers(monitor)
return routes.handleHealthServers()
},
)

Expand All @@ -109,14 +114,14 @@ func RegisterHealthRoutes(routerAPI huma.API, monitor contracts.MCPHealthMonitor
Tags: tags,
},
func(ctx context.Context, input *ServerHealthRequest) (*ServerHealthResponse, error) {
return handleHealthServer(monitor, input.Name)
return routes.handleHealthServer(input.Name)
},
)
}

// handleHealthServers is the handler for retrieving the current health for all registered MCP servers.
func handleHealthServers(monitor contracts.MCPHealthMonitor) (*ServersHealthResponse, error) {
servers := monitor.List()
func (r *healthRoutes) handleHealthServers() (*ServersHealthResponse, error) {
servers := r.monitor.List()

slices.SortFunc(servers, func(a, b domain.ServerHealth) int {
return strings.Compare(a.Name, b.Name)
Expand All @@ -138,8 +143,8 @@ func handleHealthServers(monitor contracts.MCPHealthMonitor) (*ServersHealthResp
}

// handleHealthServer is the handler for retrieving the current health the specified registered MCP server.
func handleHealthServer(monitor contracts.MCPHealthMonitor, name string) (*ServerHealthResponse, error) {
health, err := monitor.Status(name)
func (r *healthRoutes) handleHealthServer(name string) (*ServerHealthResponse, error) {
health, err := r.monitor.Status(name)
if err != nil {
return nil, err
}
Expand Down
8 changes: 6 additions & 2 deletions internal/api/health_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,10 @@ func newMockHealthMonitor() *mockHealthMonitor {
}
}

func newTestHealthRoutes(monitor *mockHealthMonitor) *healthRoutes {
return &healthRoutes{monitor: monitor}
}

func (m *mockHealthMonitor) Status(name string) (domain.ServerHealth, error) {
if health, ok := m.servers[name]; ok {
return health, nil
Expand Down Expand Up @@ -111,7 +115,7 @@ func TestHandleHealthServer_ServerNotTracked(t *testing.T) {
monitor := newMockHealthMonitor()

// Try to get health for a server that doesn't exist.
result, err := handleHealthServer(monitor, "nonexistent-server")
result, err := newTestHealthRoutes(monitor).handleHealthServer("nonexistent-server")
require.Error(t, err)
require.Nil(t, result)

Expand All @@ -129,7 +133,7 @@ func TestHandleHealthServer_ServerExists(t *testing.T) {
require.NoError(t, err)

// Get health for existing server.
result, err := handleHealthServer(monitor, "test-server")
result, err := newTestHealthRoutes(monitor).handleHealthServer("test-server")
require.NoError(t, err)
require.NotNil(t, result)

Expand Down
19 changes: 8 additions & 11 deletions internal/api/prompts.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@ import (
"github.com/danielgtaylor/huma/v2"
"github.com/mark3labs/mcp-go/mcp"

"github.com/mozilla-ai/mcpd/internal/contracts"
errorsint "github.com/mozilla-ai/mcpd/internal/errors"
)

Expand Down Expand Up @@ -144,12 +143,11 @@ func (d DomainPromptMessage) ToAPIType() (PromptMessage, error) {
}

// handleServerPrompts returns the list of prompts for a given server.
func handleServerPrompts(
accessor contracts.MCPClientAccessor,
func (r *serverRoutes) handleServerPrompts(
name string,
cursor string,
) (*PromptsListResponse, error) {
mcpClient, clientOk := accessor.Client(name)
mcpClient, clientOk := r.accessor.Client(name)
if !clientOk {
return nil, fmt.Errorf("%w: %s", errorsint.ErrServerNotFound, name)
}
Expand Down Expand Up @@ -194,13 +192,12 @@ func handleServerPrompts(
}

// handleServerPromptGenerate generates a prompt from a template on a server.
func handleServerPromptGenerate(
accessor contracts.MCPClientAccessor,
func (r *serverRoutes) handleServerPromptGenerate(
serverName string,
promptName string,
arguments map[string]string,
) (*GeneratePromptResponse, error) {
mcpClient, clientOk := accessor.Client(serverName)
mcpClient, clientOk := r.accessor.Client(serverName)
if !clientOk {
return nil, fmt.Errorf("%w: %s", errorsint.ErrServerNotFound, serverName)
}
Expand Down Expand Up @@ -240,8 +237,8 @@ func handleServerPromptGenerate(
return resp, nil
}

// RegisterPromptRoutes registers prompt-related routes under the servers API.
func RegisterPromptRoutes(parentAPI huma.API, accessor contracts.MCPClientAccessor) {
// registerPromptRoutes registers prompt-related routes under the servers API.
func (r *serverRoutes) registerPromptRoutes(parentAPI huma.API) {
tags := []string{"Prompts"}

huma.Register(
Expand All @@ -254,7 +251,7 @@ func RegisterPromptRoutes(parentAPI huma.API, accessor contracts.MCPClientAccess
Tags: tags,
},
func(ctx context.Context, input *ServerPromptsListRequest) (*PromptsListResponse, error) {
return handleServerPrompts(accessor, input.Name, input.Cursor)
return r.handleServerPrompts(input.Name, input.Cursor)
},
)

Expand All @@ -269,7 +266,7 @@ func RegisterPromptRoutes(parentAPI huma.API, accessor contracts.MCPClientAccess
Tags: tags,
},
func(ctx context.Context, input *ServerPromptGenerateRequest) (*GeneratePromptResponse, error) {
return handleServerPromptGenerate(accessor, input.ServerName, input.PromptName, input.Body.Arguments)
return r.handleServerPromptGenerate(input.ServerName, input.PromptName, input.Body.Arguments)
},
)
}
26 changes: 13 additions & 13 deletions internal/api/prompts_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -86,7 +86,7 @@ func TestAPI_HandleServerPrompts_Success(t *testing.T) {
accessor := newMockMCPClientAccessor()
accessor.Add(tc.serverName, mockClient, []string{})

result, err := handleServerPrompts(accessor, tc.serverName, "")
result, err := newTestServerRoutes(accessor).handleServerPrompts(tc.serverName, "")

require.NoError(t, err)
require.NotNil(t, result)
Expand Down Expand Up @@ -114,7 +114,7 @@ func TestAPI_HandleServerPrompts_WithCursor(t *testing.T) {
accessor := newMockMCPClientAccessor()
accessor.Add("test-server", mockClient, []string{})

result, err := handleServerPrompts(accessor, "test-server", cursor)
result, err := newTestServerRoutes(accessor).handleServerPrompts("test-server", cursor)

require.NoError(t, err)
require.NotNil(t, result)
Expand All @@ -126,7 +126,7 @@ func TestAPI_HandleServerPrompts_ServerNotFound(t *testing.T) {

accessor := newMockMCPClientAccessor()

result, err := handleServerPrompts(accessor, "nonexistent-server", "")
result, err := newTestServerRoutes(accessor).handleServerPrompts("nonexistent-server", "")

require.Error(t, err)
require.Nil(t, result)
Expand All @@ -143,7 +143,7 @@ func TestAPI_HandleServerPrompts_ListError(t *testing.T) {
accessor := newMockMCPClientAccessor()
accessor.Add("test-server", mockClient, []string{})

result, err := handleServerPrompts(accessor, "test-server", "")
result, err := newTestServerRoutes(accessor).handleServerPrompts("test-server", "")

require.Error(t, err)
require.Nil(t, result)
Expand All @@ -160,7 +160,7 @@ func TestAPI_HandleServerPrompts_NilResult(t *testing.T) {
accessor := newMockMCPClientAccessor()
accessor.Add("test-server", mockClient, []string{})

result, err := handleServerPrompts(accessor, "test-server", "")
result, err := newTestServerRoutes(accessor).handleServerPrompts("test-server", "")

require.Error(t, err)
require.Nil(t, result)
Expand All @@ -177,7 +177,7 @@ func TestAPI_HandleServerPrompts_MethodNotFound(t *testing.T) {
accessor := newMockMCPClientAccessor()
accessor.Add("test-server", mockClient, []string{})

result, err := handleServerPrompts(accessor, "test-server", "")
result, err := newTestServerRoutes(accessor).handleServerPrompts("test-server", "")

require.Error(t, err)
require.Nil(t, result)
Expand Down Expand Up @@ -205,7 +205,7 @@ func TestAPI_HandleServerPromptGenerate_Success(t *testing.T) {
promptName := "test-prompt"
arguments := map[string]string{}

result, err := handleServerPromptGenerate(accessor, "test-server", promptName, arguments)
result, err := newTestServerRoutes(accessor).handleServerPromptGenerate("test-server", promptName, arguments)

require.NoError(t, err)
require.NotNil(t, result)
Expand Down Expand Up @@ -238,7 +238,7 @@ func TestAPI_HandleServerPromptGenerate_WithArguments(t *testing.T) {
"param2": "value2",
}

result, err := handleServerPromptGenerate(accessor, "test-server", promptName, arguments)
result, err := newTestServerRoutes(accessor).handleServerPromptGenerate("test-server", promptName, arguments)

require.NoError(t, err)
require.NotNil(t, result)
Expand Down Expand Up @@ -272,7 +272,7 @@ func TestAPI_HandleServerPromptGenerate_MultipleMessages(t *testing.T) {
promptName := "multi-prompt"
arguments := map[string]string{}

result, err := handleServerPromptGenerate(accessor, "test-server", promptName, arguments)
result, err := newTestServerRoutes(accessor).handleServerPromptGenerate("test-server", promptName, arguments)

require.NoError(t, err)
require.NotNil(t, result)
Expand All @@ -289,7 +289,7 @@ func TestAPI_HandleServerPromptGenerate_ServerNotFound(t *testing.T) {
promptName := "test-prompt"
arguments := map[string]string{}

result, err := handleServerPromptGenerate(accessor, "nonexistent-server", promptName, arguments)
result, err := newTestServerRoutes(accessor).handleServerPromptGenerate("nonexistent-server", promptName, arguments)

require.Error(t, err)
require.Nil(t, result)
Expand All @@ -309,7 +309,7 @@ func TestAPI_HandleServerPromptGenerate_GenerateError(t *testing.T) {
promptName := "nonexistent-prompt"
arguments := map[string]string{}

result, err := handleServerPromptGenerate(accessor, "test-server", promptName, arguments)
result, err := newTestServerRoutes(accessor).handleServerPromptGenerate("test-server", promptName, arguments)

require.Error(t, err)
require.Nil(t, result)
Expand All @@ -329,7 +329,7 @@ func TestAPI_HandleServerPromptGenerate_NilResult(t *testing.T) {
promptName := "test-prompt"
arguments := map[string]string{}

result, err := handleServerPromptGenerate(accessor, "test-server", promptName, arguments)
result, err := newTestServerRoutes(accessor).handleServerPromptGenerate("test-server", promptName, arguments)

require.Error(t, err)
require.Nil(t, result)
Expand All @@ -349,7 +349,7 @@ func TestAPI_HandleServerPromptGenerate_MethodNotFound(t *testing.T) {
promptName := "test-prompt"
arguments := map[string]string{}

result, err := handleServerPromptGenerate(accessor, "test-server", promptName, arguments)
result, err := newTestServerRoutes(accessor).handleServerPromptGenerate("test-server", promptName, arguments)

require.Error(t, err)
require.Nil(t, result)
Expand Down
26 changes: 11 additions & 15 deletions internal/api/resources.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@ import (
"github.com/danielgtaylor/huma/v2"
"github.com/mark3labs/mcp-go/mcp"

"github.com/mozilla-ai/mcpd/internal/contracts"
errorsint "github.com/mozilla-ai/mcpd/internal/errors"
)

Expand Down Expand Up @@ -164,12 +163,11 @@ func (d DomainResourceTemplate) ToAPIType() (ResourceTemplate, error) {
}

// handleServerResources returns the list of resources for a given server.
func handleServerResources(
accessor contracts.MCPClientAccessor,
func (r *serverRoutes) handleServerResources(
name string,
cursor string,
) (*ResourcesResponse, error) {
mcpClient, clientOk := accessor.Client(name)
mcpClient, clientOk := r.accessor.Client(name)
if !clientOk {
return nil, fmt.Errorf("%w: %s", errorsint.ErrServerNotFound, name)
}
Expand Down Expand Up @@ -214,12 +212,11 @@ func handleServerResources(
}

// handleServerResourceTemplates returns the list of resource templates for a given server.
func handleServerResourceTemplates(
accessor contracts.MCPClientAccessor,
func (r *serverRoutes) handleServerResourceTemplates(
name string,
cursor string,
) (*ResourceTemplatesResponse, error) {
mcpClient, clientOk := accessor.Client(name)
mcpClient, clientOk := r.accessor.Client(name)
if !clientOk {
return nil, fmt.Errorf("%w: %s", errorsint.ErrServerNotFound, name)
}
Expand Down Expand Up @@ -264,12 +261,11 @@ func handleServerResourceTemplates(
}

// handleServerResourceContent gets the content of a specific resource from a server.
func handleServerResourceContent(
accessor contracts.MCPClientAccessor,
func (r *serverRoutes) handleServerResourceContent(
name string,
uri string,
) (*ResourceContentResponse, error) {
mcpClient, clientOk := accessor.Client(name)
mcpClient, clientOk := r.accessor.Client(name)
if !clientOk {
return nil, fmt.Errorf("%w: %s", errorsint.ErrServerNotFound, name)
}
Expand Down Expand Up @@ -318,8 +314,8 @@ func handleServerResourceContent(
return resp, nil
}

// RegisterResourceRoutes registers resource-related routes under the servers API.
func RegisterResourceRoutes(parentAPI huma.API, accessor contracts.MCPClientAccessor) {
// registerResourceRoutes registers resource-related routes under the servers API.
func (r *serverRoutes) registerResourceRoutes(parentAPI huma.API) {
tags := []string{"Resources"}

huma.Register(
Expand All @@ -332,7 +328,7 @@ func RegisterResourceRoutes(parentAPI huma.API, accessor contracts.MCPClientAcce
Tags: tags,
},
func(ctx context.Context, input *ServerResourcesRequest) (*ResourcesResponse, error) {
return handleServerResources(accessor, input.Name, input.Cursor)
return r.handleServerResources(input.Name, input.Cursor)
},
)

Expand All @@ -346,7 +342,7 @@ func RegisterResourceRoutes(parentAPI huma.API, accessor contracts.MCPClientAcce
Tags: tags,
},
func(ctx context.Context, input *ServerResourceTemplatesRequest) (*ResourceTemplatesResponse, error) {
return handleServerResourceTemplates(accessor, input.Name, input.Cursor)
return r.handleServerResourceTemplates(input.Name, input.Cursor)
},
)

Expand All @@ -361,7 +357,7 @@ func RegisterResourceRoutes(parentAPI huma.API, accessor contracts.MCPClientAcce
Tags: tags,
},
func(ctx context.Context, input *ServerResourceContentRequest) (*ResourceContentResponse, error) {
return handleServerResourceContent(accessor, input.Name, input.URI)
return r.handleServerResourceContent(input.Name, input.URI)
},
)
}
Loading
Loading