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
11 changes: 4 additions & 7 deletions client/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -742,14 +742,11 @@ func (c *Client) handleSamplingRequestTransport(ctx context.Context, request tra
// Fix content parsing - HTTP transport unmarshals TextContent as map[string]any
// Use the helper function to properly handle content from different transports
for i := range params.Messages {
if contentMap, ok := params.Messages[i].Content.(map[string]any); ok {
// Parse the content map into a proper Content type
content, err := mcp.ParseContent(contentMap)
if err != nil {
return nil, fmt.Errorf("failed to parse content for message %d: %w", i, err)
}
params.Messages[i].Content = content
content, err := mcp.ParseSamplingContent(params.Messages[i].Content)
if err != nil {
return nil, fmt.Errorf("failed to parse content for message %d: %w", i, err)
}
params.Messages[i].Content = content
}

// Create the MCP request
Expand Down
66 changes: 64 additions & 2 deletions client/sampling_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,17 +5,22 @@ import (
"encoding/json"
"testing"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"

"github.com/mark3labs/mcp-go/client/transport"
"github.com/mark3labs/mcp-go/mcp"
)

// mockSamplingHandler implements SamplingHandler for testing
type mockSamplingHandler struct {
result *mcp.CreateMessageResult
err error
result *mcp.CreateMessageResult
err error
request mcp.CreateMessageRequest
}

func (m *mockSamplingHandler) CreateMessage(ctx context.Context, request mcp.CreateMessageRequest) (*mcp.CreateMessageResult, error) {
m.request = request
if m.err != nil {
return nil, m.err
}
Expand Down Expand Up @@ -87,6 +92,63 @@ func TestClient_HandleSamplingRequest(t *testing.T) {
}
}

func TestClient_HandleSamplingRequestArrayContent(t *testing.T) {
handler := &mockSamplingHandler{
result: &mcp.CreateMessageResult{
SamplingMessage: mcp.SamplingMessage{
Role: mcp.RoleAssistant,
Content: mcp.NewTextContent("Paris is warmer than London."),
},
Model: "test-model",
},
}
client := &Client{samplingHandler: handler}

params := json.RawMessage(`{
"messages": [
{"role": "user", "content": {"type": "text", "text": "What's the weather like in Paris and London?"}},
{"role": "assistant", "content": [
{"type": "tool_use", "id": "call_abc123", "name": "get_weather", "input": {"city": "Paris"}},
{"type": "tool_use", "id": "call_def456", "name": "get_weather", "input": {"city": "London"}}
]},
{"role": "user", "content": [
{"type": "tool_result", "toolUseId": "call_abc123", "content": [{"type": "text", "text": "18°C, partly cloudy"}]},
{"type": "tool_result", "toolUseId": "call_def456", "content": [{"type": "text", "text": "15°C, rainy"}]}
]}
],
"maxTokens": 1000
}`)

_, err := client.handleIncomingRequest(t.Context(), transport.JSONRPCRequest{
JSONRPC: mcp.JSONRPC_VERSION,
ID: mcp.NewRequestId(1),
Method: string(mcp.MethodSamplingCreateMessage),
Params: params,
})
require.NoError(t, err)

assert.Equal(t, []mcp.SamplingMessage{
{
Role: mcp.RoleUser,
Content: mcp.NewTextContent("What's the weather like in Paris and London?"),
},
{
Role: mcp.RoleAssistant,
Content: []mcp.Content{
mcp.NewToolUseContent("call_abc123", "get_weather", map[string]any{"city": "Paris"}),
mcp.NewToolUseContent("call_def456", "get_weather", map[string]any{"city": "London"}),
},
},
{
Role: mcp.RoleUser,
Content: []mcp.Content{
mcp.NewToolResultContent("call_abc123", []mcp.Content{mcp.NewTextContent("18°C, partly cloudy")}, false),
mcp.NewToolResultContent("call_def456", []mcp.Content{mcp.NewTextContent("15°C, rainy")}, false),
},
},
}, handler.request.Messages)
}

func TestWithSamplingHandler(t *testing.T) {
handler := &mockSamplingHandler{}
client := &Client{}
Expand Down
2 changes: 1 addition & 1 deletion mcp/types.go
Original file line number Diff line number Diff line change
Expand Up @@ -1209,7 +1209,7 @@ type CreateMessageResult struct {
// SamplingMessage describes a message issued to or received from an LLM API.
type SamplingMessage struct {
Role Role `json:"role"`
Content any `json:"content"` // Can be TextContent, ImageContent, AudioContent, ToolUseContent or ToolResultContent
Content any `json:"content"` // Can be TextContent, ImageContent, AudioContent, ToolUseContent or ToolResultContent, or a []Content of them
}

type Annotations struct {
Expand Down
27 changes: 27 additions & 0 deletions mcp/utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -730,6 +730,33 @@ func ParseContent(contentMap map[string]any) (Content, error) {
return nil, fmt.Errorf("unsupported content type: %s", contentType)
}

// ParseSamplingContent parses the content of a SamplingMessage that was
// decoded into generic JSON values. A single content block is returned as a
// Content and an array of blocks as a []Content. Any other value, such as
// content that is already typed, is returned unchanged.
func ParseSamplingContent(content any) (any, error) {
switch c := content.(type) {
case map[string]any:
return ParseContent(c)
case []any:
blocks := make([]Content, 0, len(c))
for i, item := range c {
itemMap, ok := item.(map[string]any)
if !ok {
return nil, fmt.Errorf("content[%d]: expected object, got %T", i, item)
}
block, err := ParseContent(itemMap)
if err != nil {
return nil, fmt.Errorf("content[%d]: %w", i, err)
}
blocks = append(blocks, block)
}
return blocks, nil
default:
return content, nil
}
}

// resultEnvelope holds the fields a result carries outside its payload: the
// SEP-2322 round-trip fields and the SEP-2549 cache hints.
type resultEnvelope struct {
Expand Down
53 changes: 53 additions & 0 deletions mcp/utils_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -713,6 +713,59 @@ func TestParseContent(t *testing.T) {
}
}

func TestParseSamplingContent(t *testing.T) {
tests := []struct {
name string
content any
expected any
expectError string
}{
{
name: "single block",
content: map[string]any{"type": "text", "text": "Hello"},
expected: NewTextContent("Hello"),
},
{
name: "array of blocks",
content: []any{
map[string]any{"type": "tool_use", "id": "call_1", "name": "get_weather", "input": map[string]any{"city": "Paris"}},
map[string]any{"type": "tool_use", "id": "call_2", "name": "get_weather", "input": map[string]any{"city": "London"}},
},
expected: []Content{
NewToolUseContent("call_1", "get_weather", map[string]any{"city": "Paris"}),
NewToolUseContent("call_2", "get_weather", map[string]any{"city": "London"}),
},
},
{
name: "typed content",
content: NewTextContent("Hello"),
expected: NewTextContent("Hello"),
},
{
name: "array entry that is not an object",
content: []any{"Hello"},
expectError: "content[0]: expected object, got string",
},
{
name: "array entry with an unknown type",
content: []any{map[string]any{"type": "video"}},
expectError: "content[0]: unsupported content type: video",
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result, err := ParseSamplingContent(tt.content)
if tt.expectError != "" {
require.EqualError(t, err, tt.expectError)
return
}
require.NoError(t, err)
assert.Equal(t, tt.expected, result)
})
}
}

func TestNewJSONRPCResultResponse(t *testing.T) {
t.Parallel()

Expand Down
12 changes: 4 additions & 8 deletions server/stdio.go
Original file line number Diff line number Diff line change
Expand Up @@ -706,15 +706,11 @@ func (s *stdioSession) handleSamplingResponse(rawMessage json.RawMessage) bool {
samplingResp.err = fmt.Errorf("failed to unmarshal sampling response: %w", err)
} else {
// Parse content from map[string]any to proper Content type (TextContent, ImageContent, AudioContent)
if contentMap, ok := result.Content.(map[string]any); ok {
content, err := mcp.ParseContent(contentMap)
if err != nil {
samplingResp.err = fmt.Errorf("failed to parse sampling response content: %w", err)
} else {
result.Content = content
samplingResp.result = &result
}
content, err := mcp.ParseSamplingContent(result.Content)
if err != nil {
samplingResp.err = fmt.Errorf("failed to parse sampling response content: %w", err)
} else {
result.Content = content
samplingResp.result = &result
}
}
Expand Down
24 changes: 24 additions & 0 deletions server/stdio_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -639,3 +639,27 @@ func TestStdioJoinsRequestHandlersOnEOF(t *testing.T) {
t.Fatal("Listen did not return once the request handler finished")
}
}

func TestStdioSessionSamplingResponseArrayContent(t *testing.T) {
session := &stdioSession{pendingRequests: make(map[int64]chan *samplingResponse)}
responseChan := make(chan *samplingResponse, 1)
session.pendingRequests[1] = responseChan

handled := session.handleSamplingResponse(json.RawMessage(`{"jsonrpc": "2.0", "id": 1, "result": {
"role": "assistant",
"content": [
{"type": "tool_use", "id": "call_abc123", "name": "get_weather", "input": {"city": "Paris"}},
{"type": "tool_use", "id": "call_def456", "name": "get_weather", "input": {"city": "London"}}
],
"model": "test-model",
"stopReason": "toolUse"
}}`))
require.True(t, handled)

response := <-responseChan
require.NoError(t, response.err)
require.Equal(t, []mcp.Content{
mcp.NewToolUseContent("call_abc123", "get_weather", map[string]any{"city": "Paris"}),
mcp.NewToolUseContent("call_def456", "get_weather", map[string]any{"city": "London"}),
}, response.result.Content)
}
10 changes: 4 additions & 6 deletions server/streamable_http.go
Original file line number Diff line number Diff line change
Expand Up @@ -1868,13 +1868,11 @@ func (s *streamableHttpSession) RequestSampling(ctx context.Context, request mcp

// Parse content from map[string]any to proper Content type (TextContent, ImageContent, AudioContent)
// HTTP transport unmarshals Content as map[string]any, we need to convert it to the proper type
if contentMap, ok := result.Content.(map[string]any); ok {
content, err := mcp.ParseContent(contentMap)
if err != nil {
return nil, fmt.Errorf("failed to parse sampling response content: %w", err)
}
result.Content = content
content, err := mcp.ParseSamplingContent(result.Content)
if err != nil {
return nil, fmt.Errorf("failed to parse sampling response content: %w", err)
}
result.Content = content

return &result, nil
case <-ctx.Done():
Expand Down
36 changes: 36 additions & 0 deletions server/streamable_http_sampling_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@ import (
"testing"
"time"

"github.com/stretchr/testify/require"

"github.com/mark3labs/mcp-go/mcp"
)

Expand Down Expand Up @@ -215,3 +217,37 @@ func TestStreamableHTTPServer_SamplingQueueFull(t *testing.T) {
t.Errorf("Expected queue full error, got: %v", err)
}
}

func TestStreamableHTTPServer_SamplingArrayContent(t *testing.T) {
session := newStreamableHttpSession("test-session", nil, nil, nil, nil, new(atomic.Int64))

go func() {
item := <-session.samplingRequestChan
item.response <- samplingResponseItem{
requestID: item.requestID,
result: json.RawMessage(`{
"role": "assistant",
"content": [
{"type": "tool_use", "id": "call_abc123", "name": "get_weather", "input": {"city": "Paris"}},
{"type": "tool_use", "id": "call_def456", "name": "get_weather", "input": {"city": "London"}}
],
"model": "test-model",
"stopReason": "toolUse"
}`),
}
}()

result, err := session.RequestSampling(t.Context(), mcp.CreateMessageRequest{
CreateMessageParams: mcp.CreateMessageParams{
Messages: []mcp.SamplingMessage{
{Role: mcp.RoleUser, Content: mcp.NewTextContent("What's the weather like in Paris and London?")},
},
MaxTokens: 1000,
},
})
require.NoError(t, err)
require.Equal(t, []mcp.Content{
mcp.NewToolUseContent("call_abc123", "get_weather", map[string]any{"city": "Paris"}),
mcp.NewToolUseContent("call_def456", "get_weather", map[string]any{"city": "London"}),
}, result.Content)
}
Loading