Skip to content
Merged
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
20 changes: 14 additions & 6 deletions internal/mcp/connection.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"fmt"
"io"
"log"
"log/slog"
"net/http"
"os"
"os/exec"
Expand Down Expand Up @@ -158,11 +159,18 @@ type Connection struct {
}

// newMCPClient creates a new MCP SDK client with standard implementation details
func newMCPClient() *sdk.Client {
// Pass nil for logger parameter to disable SDK logging (for tests)
func newMCPClient(log *logger.Logger) *sdk.Client {
var slogLogger *slog.Logger
if log != nil {
slogLogger = logger.NewSlogLoggerWithHandler(log)
}
return sdk.NewClient(&sdk.Implementation{
Name: "awmg",
Version: version.Get(),
}, &sdk.ClientOptions{})
}, &sdk.ClientOptions{
Logger: slogLogger,
})
}

// newHTTPConnection creates a new HTTP Connection struct with common fields
Expand Down Expand Up @@ -300,8 +308,8 @@ func NewConnection(ctx context.Context, serverID, command string, args []string,
logConn.Printf("Creating new MCP connection: command=%s, args=%v", command, sanitize.SanitizeArgs(args))
ctx, cancel := context.WithCancel(ctx)

// Create MCP client
client := newMCPClient()
// Create MCP client with logger
client := newMCPClient(logConn)

// Expand Docker -e flags that reference environment variables
// Docker's `-e VAR_NAME` expects VAR_NAME to be in the environment
Expand Down Expand Up @@ -506,8 +514,8 @@ func trySDKTransport(
transportName string,
createTransport transportConnector,
) (*Connection, error) {
// Create MCP client
client := newMCPClient()
// Create MCP client with logger
client := newMCPClient(logConn)

// Create transport using the provided connector
transport := createTransport(url, httpClient)
Expand Down
12 changes: 10 additions & 2 deletions internal/mcp/connection_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (
"strings"
"testing"

"github.com/github/gh-aw-mcpg/internal/logger"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
Expand Down Expand Up @@ -568,10 +569,17 @@ data: {"jsonrpc":"2.0","id":` + idStr + `,"result":{"tools":[]}}

// TestNewMCPClient tests the newMCPClient helper function
func TestNewMCPClient(t *testing.T) {
client := newMCPClient()
client := newMCPClient(nil)
require.NotNil(t, client, "newMCPClient should return a non-nil client")
}

// TestNewMCPClientWithLogger tests that newMCPClient accepts a logger
func TestNewMCPClientWithLogger(t *testing.T) {
log := logger.New("test:client")
client := newMCPClient(log)
require.NotNil(t, client, "newMCPClient should return a non-nil client with logger")
}

// TestCreateJSONRPCRequest tests the createJSONRPCRequest helper function
func TestCreateJSONRPCRequest(t *testing.T) {
tests := []struct {
Expand Down Expand Up @@ -687,7 +695,7 @@ func TestNewHTTPConnection(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()

client := newMCPClient()
client := newMCPClient(nil)
url := "http://example.com/mcp"
headers := map[string]string{"Authorization": "test"}
httpClient := &http.Client{}
Expand Down
6 changes: 4 additions & 2 deletions internal/server/routed.go
Original file line number Diff line number Diff line change
Expand Up @@ -183,11 +183,13 @@ func CreateHTTPServerForRoutedMode(addr string, unifiedServer *UnifiedServer, ap
func createFilteredServer(unifiedServer *UnifiedServer, backendID string) *sdk.Server {
logRouted.Printf("Creating filtered server: backend=%s", backendID)

// Create a new SDK server for this route
// Create a new SDK server for this route with logger
server := sdk.NewServer(&sdk.Implementation{
Name: fmt.Sprintf("awmg-%s", backendID),
Version: "1.0.0",
}, nil)
}, &sdk.ServerOptions{
Logger: logger.NewSlogLoggerWithHandler(logRouted),
})

// Get tools for this backend from the unified server
tools := unifiedServer.GetToolsForBackend(backendID)
Expand Down
6 changes: 4 additions & 2 deletions internal/server/unified.go
Original file line number Diff line number Diff line change
Expand Up @@ -132,11 +132,13 @@ func NewUnified(ctx context.Context, cfg *config.Config) (*UnifiedServer, error)
enableDIFC: cfg.EnableDIFC,
}

// Create MCP server
// Create MCP server with logger
server := sdk.NewServer(&sdk.Implementation{
Name: "awmg-unified",
Version: "1.0.0",
}, nil)
}, &sdk.ServerOptions{
Logger: logger.NewSlogLoggerWithHandler(logUnified),
})

us.server = server

Expand Down