client: Honor IsError and extract JSON from tool text
CI / Test and build (pull_request) Successful in 13s

#15
This commit is contained in:
ash
2026-09-05 04:53:47 +00:00
parent 11feaad0e3
commit 66ec6305e8
2 changed files with 206 additions and 9 deletions
+128 -7
View File
@@ -5,7 +5,9 @@ import (
"encoding/json"
"fmt"
"net/http"
"strings"
"time"
"unicode/utf8"
"github.com/modelcontextprotocol/go-sdk/auth"
mcp "github.com/modelcontextprotocol/go-sdk/mcp"
@@ -98,9 +100,19 @@ func toolJSON(res *mcp.CallToolResult) (json.RawMessage, error) {
if res == nil {
return nil, fmt.Errorf("empty tool result")
}
// GetError is only set by SetError on the server; the err field is not marshaled to clients.
if err := res.GetError(); err != nil {
return nil, err
}
texts, types := collectTextAndTypes(res)
concat := strings.Join(texts, "")
// IsError is the client-visible soft-fail flag; error text lives in Content.
if res.IsError {
if concat == "" {
return nil, fmt.Errorf("tool error")
}
return nil, fmt.Errorf("%s", concat)
}
if res.StructuredContent != nil {
b, err := json.Marshal(res.StructuredContent)
if err != nil {
@@ -108,21 +120,130 @@ func toolJSON(res *mcp.CallToolResult) (json.RawMessage, error) {
}
return json.RawMessage(b), nil
}
var b []byte
for _, t := range texts {
if json.Valid([]byte(t)) {
return json.RawMessage(t), nil
}
}
if concat == "" {
return json.RawMessage(`{}`), nil
}
if json.Valid([]byte(concat)) {
return json.RawMessage(concat), nil
}
if raw := extractFirstJSON(concat); raw != nil {
return raw, nil
}
structured := "nil"
if res.StructuredContent != nil {
structured = "present"
}
return nil, fmt.Errorf("non-json tool result (isError=%v types=[%s] structured=%s textLen=%d prefix=%q)",
res.IsError, strings.Join(types, " "), structured, utf8.RuneCountInString(concat), runePrefix(concat, 80))
}
func collectTextAndTypes(res *mcp.CallToolResult) (texts, types []string) {
for _, c := range res.Content {
types = append(types, contentTypeName(c))
t, ok := c.(*mcp.TextContent)
if !ok {
continue
}
b = append(b, t.Text...)
texts = append(texts, t.Text)
}
if len(b) == 0 {
return json.RawMessage(`{}`), nil
return texts, types
}
func contentTypeName(c mcp.Content) string {
switch c.(type) {
case *mcp.TextContent:
return "text"
case *mcp.ImageContent:
return "image"
case *mcp.AudioContent:
return "audio"
case *mcp.ResourceLink:
return "resource_link"
case *mcp.EmbeddedResource:
return "resource"
case *mcp.ToolUseContent:
return "tool_use"
case *mcp.ToolResultContent:
return "tool_result"
default:
return "unknown"
}
if !json.Valid(b) {
return nil, fmt.Errorf("non-json tool result")
}
func runePrefix(s string, n int) string {
if n <= 0 || s == "" {
return ""
}
return json.RawMessage(b), nil
i := 0
for j := range s {
if i == n {
return s[:j]
}
i++
}
return s
}
// extractFirstJSON returns the first balanced {...} or [...] substring that is
// valid JSON. Brace matching skips quoted strings and respects escapes.
func extractFirstJSON(s string) json.RawMessage {
for i := 0; i < len(s); i++ {
if s[i] != '{' && s[i] != '[' {
continue
}
end := balancedJSONEnd(s, i)
if end <= i {
continue
}
cand := s[i:end]
if json.Valid([]byte(cand)) {
return json.RawMessage(cand)
}
}
return nil
}
func balancedJSONEnd(s string, start int) int {
depth := 0
inString := false
escape := false
for i := start; i < len(s); i++ {
c := s[i]
if inString {
if escape {
escape = false
continue
}
if c == '\\' {
escape = true
continue
}
if c == '"' {
inString = false
}
continue
}
switch c {
case '"':
inString = true
case '{', '[':
depth++
case '}', ']':
depth--
if depth == 0 {
return i + 1
}
if depth < 0 {
return -1
}
}
}
return -1
}
type bearerRT struct {
+78 -2
View File
@@ -3,6 +3,7 @@ package client
import (
"encoding/json"
"errors"
"strings"
"testing"
mcp "github.com/modelcontextprotocol/go-sdk/mcp"
@@ -42,6 +43,37 @@ func TestToolJSON(t *testing.T) {
res: &mcp.CallToolResult{},
want: `{}`,
},
{
name: "wrapped JSON in prose",
res: &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: `here is data: {"symbol":"MU"} thanks`}},
},
want: `{"symbol":"MU"}`,
},
{
name: "one TextContent is JSON",
res: &mcp.CallToolResult{
Content: []mcp.Content{
&mcp.TextContent{Text: "note: "},
&mcp.TextContent{Text: `{"symbol":"MU"}`},
},
},
want: `{"symbol":"MU"}`,
},
{
name: "wrapped JSON array",
res: &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: `items: [1,{"a":2}] done`}},
},
want: `[1,{"a":2}]`,
},
{
name: "wrapped JSON with brace in string",
res: &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: `here {"msg":"say } hi"} x`}},
},
want: `{"msg":"say } hi"}`,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
@@ -76,12 +108,56 @@ func TestToolJSON_getError(t *testing.T) {
}
}
func TestToolJSON_isErrorText(t *testing.T) {
t.Parallel()
res := &mcp.CallToolResult{
IsError: true,
Content: []mcp.Content{&mcp.TextContent{Text: "historicals unavailable"}},
}
if res.GetError() != nil {
t.Fatal("GetError should be nil when IsError is set via field")
}
_, err := toolJSON(res)
if err == nil {
t.Fatal("expected error")
}
if err.Error() != "historicals unavailable" {
t.Fatalf("%v", err)
}
if strings.Contains(err.Error(), "non-json tool result") {
t.Fatalf("got non-json wrapping: %v", err)
}
}
func TestToolJSON_isErrorEmpty(t *testing.T) {
t.Parallel()
_, err := toolJSON(&mcp.CallToolResult{IsError: true})
if err == nil {
t.Fatal("expected error")
}
if err.Error() != "tool error" {
t.Fatalf("%v", err)
}
if strings.Contains(err.Error(), "non-json tool result") {
t.Fatalf("got non-json wrapping: %v", err)
}
}
func TestToolJSON_nonJSONText(t *testing.T) {
t.Parallel()
_, err := toolJSON(&mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: "not json"}},
})
if err == nil || err.Error() != "non-json tool result" {
t.Fatalf("%v", err)
if err == nil {
t.Fatal("expected error")
}
msg := err.Error()
if !strings.Contains(msg, "non-json") {
t.Fatalf("missing non-json: %v", err)
}
for _, want := range []string{`isError=false`, `types=[text]`, `structured=nil`, `prefix="not json"`} {
if !strings.Contains(msg, want) {
t.Fatalf("missing %q in %v", want, err)
}
}
}