Files
ai-agent/internal/ai/mcps/client.go
T

201 lines
4.7 KiB
Go
Raw Normal View History

2026-04-09 10:01:23 +08:00
package mcps
import (
"context"
"net/http"
"strings"
"time"
"agent-desk/internal/pkg/errorsx"
"agent-desk/internal/pkg/i18nx"
2026-04-09 10:01:23 +08:00
"github.com/modelcontextprotocol/go-sdk/mcp"
)
type Client struct{}
func NewClient() *Client {
return &Client{}
}
func (c *Client) TestConnection(ctx context.Context, cfg ServerConfig) (*ConnectionResult, error) {
session, closeFn, err := c.connect(ctx, cfg)
if err != nil {
return nil, err
}
defer closeFn()
initResult := session.InitializeResult()
serverName := ""
version := ""
protocol := ""
if initResult != nil {
serverName = initResult.ServerInfo.Name
version = initResult.ServerInfo.Version
protocol = initResult.ProtocolVersion
}
return &ConnectionResult{
ServerCode: cfg.Code,
Endpoint: cfg.Endpoint,
Protocol: protocol,
ServerName: serverName,
Version: version,
}, nil
}
func (c *Client) ListTools(ctx context.Context, cfg ServerConfig) ([]ToolInfo, error) {
session, closeFn, err := c.connect(ctx, cfg)
if err != nil {
return nil, err
}
defer closeFn()
result, err := session.ListTools(ctx, nil)
if err != nil {
return nil, i18nx.Errorf("error.mcp.listToolsFailed", err)
2026-04-09 10:01:23 +08:00
}
ret := make([]ToolInfo, 0, len(result.Tools))
for _, tool := range result.Tools {
readOnlyHint := tool.Annotations != nil && tool.Annotations.ReadOnlyHint
2026-04-09 10:01:23 +08:00
ret = append(ret, ToolInfo{
Name: tool.Name,
Title: tool.Title,
Description: tool.Description,
InputSchema: tool.InputSchema,
OutputSchema: tool.OutputSchema,
ReadOnlyHint: readOnlyHint,
2026-04-09 10:01:23 +08:00
})
}
return ret, nil
}
func (c *Client) CallTool(ctx context.Context, cfg ServerConfig, toolName string, arguments map[string]any) (*ToolCallResult, error) {
toolName = strings.TrimSpace(toolName)
if toolName == "" {
return nil, errorsx.InvalidParamI18n("error.e0076")
2026-04-09 10:01:23 +08:00
}
session, closeFn, err := c.connect(ctx, cfg)
if err != nil {
return nil, err
}
defer closeFn()
result, err := session.CallTool(ctx, &mcp.CallToolParams{
Name: toolName,
Arguments: arguments,
})
if err != nil {
return nil, i18nx.Errorf("error.mcp.callToolFailed", err)
2026-04-09 10:01:23 +08:00
}
return &ToolCallResult{
ServerCode: cfg.Code,
ToolName: toolName,
IsError: result.IsError,
Content: convertContents(result.Content),
StructuredContent: result.StructuredContent,
}, nil
}
func (c *Client) connect(ctx context.Context, cfg ServerConfig) (*mcp.ClientSession, func(), error) {
if strings.TrimSpace(cfg.Code) == "" {
return nil, nil, errorsx.InvalidParamI18n("error.e0070")
2026-04-09 10:01:23 +08:00
}
if strings.TrimSpace(cfg.Endpoint) == "" {
return nil, nil, errorsx.InvalidParamI18n("error.e0032")
2026-04-09 10:01:23 +08:00
}
timeout := time.Duration(cfg.TimeoutMS) * time.Millisecond
if timeout <= 0 {
timeout = 15 * time.Second
}
connCtx, cancel := context.WithTimeout(ctx, timeout)
httpClient := &http.Client{
Transport: &headerRoundTripper{
next: http.DefaultTransport,
headers: cfg.Headers,
},
}
client := mcp.NewClient(&mcp.Implementation{
Name: "agent-desk-mcp-client",
2026-04-09 10:01:23 +08:00
Version: "v1",
}, nil)
transport := &mcp.StreamableClientTransport{
Endpoint: cfg.Endpoint,
HTTPClient: httpClient,
MaxRetries: 0,
DisableStandaloneSSE: true,
}
session, err := client.Connect(connCtx, transport, nil)
if err != nil {
cancel()
return nil, nil, i18nx.Errorf("error.mcp.connectServerFailed", err)
2026-04-09 10:01:23 +08:00
}
return session, func() {
_ = session.Close()
cancel()
}, nil
}
func convertContents(contents []mcp.Content) []ToolResultContent {
ret := make([]ToolResultContent, 0, len(contents))
for _, item := range contents {
switch v := item.(type) {
case *mcp.TextContent:
ret = append(ret, ToolResultContent{
Type: "text",
Text: v.Text,
})
case *mcp.ImageContent:
ret = append(ret, ToolResultContent{
Type: "image",
Data: map[string]any{
"mimeType": v.MIMEType,
"data": v.Data,
},
})
case *mcp.AudioContent:
ret = append(ret, ToolResultContent{
Type: "audio",
Data: map[string]any{
"mimeType": v.MIMEType,
"data": v.Data,
},
})
case *mcp.EmbeddedResource:
ret = append(ret, ToolResultContent{
Type: "resource",
Data: v.Resource,
})
default:
ret = append(ret, ToolResultContent{
Type: "unknown",
Data: v,
})
}
}
return ret
}
type headerRoundTripper struct {
next http.RoundTripper
headers map[string]string
}
func (r *headerRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
next := r.next
if next == nil {
next = http.DefaultTransport
}
clone := req.Clone(req.Context())
for key, value := range r.headers {
key = strings.TrimSpace(key)
if key == "" {
continue
}
clone.Header.Set(key, value)
}
return next.RoundTrip(clone)
}