go-common中增加mcp抽离成通用方法(编排由业务方自行处理)
This commit is contained in:
200
mcp/manager_test.go
Normal file
200
mcp/manager_test.go
Normal file
@@ -0,0 +1,200 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"git.toowon.com/jimmy/go-common/config"
|
||||
)
|
||||
|
||||
// fakeClient 是测试用的 Client 假实现,不依赖真实网络连接。
|
||||
type fakeClient struct {
|
||||
name string
|
||||
tools []ToolDescriptor
|
||||
listErr error
|
||||
callResult ToolCallResult
|
||||
callErr error
|
||||
closed bool
|
||||
|
||||
lastCalledTool string
|
||||
lastCalledArgs map[string]any
|
||||
}
|
||||
|
||||
func (f *fakeClient) Name() string { return f.name }
|
||||
|
||||
func (f *fakeClient) ListTools(ctx context.Context) ([]ToolDescriptor, error) {
|
||||
if f.listErr != nil {
|
||||
return nil, f.listErr
|
||||
}
|
||||
return f.tools, nil
|
||||
}
|
||||
|
||||
func (f *fakeClient) CallTool(ctx context.Context, name string, args map[string]any) (ToolCallResult, error) {
|
||||
f.lastCalledTool = name
|
||||
f.lastCalledArgs = args
|
||||
if f.callErr != nil {
|
||||
return ToolCallResult{}, f.callErr
|
||||
}
|
||||
return f.callResult, nil
|
||||
}
|
||||
|
||||
func (f *fakeClient) Close() error {
|
||||
f.closed = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func newTestManager(entries ...*serverEntry) *manager {
|
||||
return &manager{enabled: true, servers: entries, routes: make(map[string]route)}
|
||||
}
|
||||
|
||||
func TestManagerDisabledWithoutConfig(t *testing.T) {
|
||||
m := NewManager(nil, nil)
|
||||
if m.Enabled() {
|
||||
t.Fatal("Manager should be disabled when cfg is nil")
|
||||
}
|
||||
if _, err := m.CallTool(context.Background(), "x__y", nil); err == nil {
|
||||
t.Fatal("CallTool should error when disabled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerDisabledWhenNoServers(t *testing.T) {
|
||||
m := NewManager(nil, &config.MCPConfig{Enabled: true})
|
||||
if m.Enabled() {
|
||||
t.Fatal("Manager should be disabled when no servers configured")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerListToolsNamespacesAndFiltersAllowlist(t *testing.T) {
|
||||
fc := &fakeClient{name: "word", tools: []ToolDescriptor{
|
||||
{Name: "create_document", Description: "create a doc"},
|
||||
{Name: "delete_document", Description: "delete a doc"},
|
||||
}}
|
||||
entry := newServerEntry(config.MCPServerConfig{
|
||||
Name: "word",
|
||||
ToolNamePrefix: "word",
|
||||
AllowedTools: []string{"create_document"},
|
||||
}, fc)
|
||||
m := newTestManager(entry)
|
||||
|
||||
tools, err := m.ListTools(context.Background())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(tools) != 1 {
|
||||
t.Fatalf("expected 1 tool after allowlist filter, got %d", len(tools))
|
||||
}
|
||||
if tools[0].FQName != "word__create_document" {
|
||||
t.Fatalf("unexpected FQName: %s", tools[0].FQName)
|
||||
}
|
||||
if tools[0].ServerName != "word" {
|
||||
t.Fatalf("unexpected ServerName: %s", tools[0].ServerName)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerListToolsSkipsFailingServer(t *testing.T) {
|
||||
ok := &fakeClient{name: "ok-server", tools: []ToolDescriptor{{Name: "t1"}}}
|
||||
bad := &fakeClient{name: "bad-server", listErr: errors.New("boom")}
|
||||
m := newTestManager(
|
||||
newServerEntry(config.MCPServerConfig{Name: "ok-server"}, ok),
|
||||
newServerEntry(config.MCPServerConfig{Name: "bad-server"}, bad),
|
||||
)
|
||||
|
||||
tools, err := m.ListTools(context.Background())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(tools) != 1 {
|
||||
t.Fatalf("expected 1 tool from healthy server, got %d", len(tools))
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCallToolRoutesByFQNameWithoutPriorListTools(t *testing.T) {
|
||||
fc := &fakeClient{name: "fetch-web", callResult: ToolCallResult{Text: "hello"}}
|
||||
entry := newServerEntry(config.MCPServerConfig{Name: "fetch-web", ToolNamePrefix: "fetch-web"}, fc)
|
||||
m := newTestManager(entry)
|
||||
|
||||
res, err := m.CallTool(context.Background(), "fetch-web__fetch_url", map[string]any{"url": "https://example.com"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.Text != "hello" {
|
||||
t.Fatalf("unexpected result text: %s", res.Text)
|
||||
}
|
||||
if fc.lastCalledTool != "fetch_url" {
|
||||
t.Fatalf("expected raw tool name %q, got %q", "fetch_url", fc.lastCalledTool)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCallToolRejectsDisallowedTool(t *testing.T) {
|
||||
entry := newServerEntry(config.MCPServerConfig{
|
||||
Name: "word",
|
||||
ToolNamePrefix: "word",
|
||||
AllowedTools: []string{"create_document"},
|
||||
}, &fakeClient{name: "word"})
|
||||
m := newTestManager(entry)
|
||||
|
||||
if _, err := m.CallTool(context.Background(), "word__delete_document", nil); err == nil {
|
||||
t.Fatal("expected error for tool not in allowlist")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCallToolUnknownServerPrefix(t *testing.T) {
|
||||
entry := newServerEntry(config.MCPServerConfig{Name: "word", ToolNamePrefix: "word"}, &fakeClient{name: "word"})
|
||||
m := newTestManager(entry)
|
||||
|
||||
if _, err := m.CallTool(context.Background(), "unknown__tool", nil); err == nil {
|
||||
t.Fatal("expected error for unresolvable fq name")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCloseClosesAllClients(t *testing.T) {
|
||||
fc1 := &fakeClient{name: "a"}
|
||||
fc2 := &fakeClient{name: "b"}
|
||||
m := newTestManager(
|
||||
newServerEntry(config.MCPServerConfig{Name: "a"}, fc1),
|
||||
newServerEntry(config.MCPServerConfig{Name: "b"}, fc2),
|
||||
)
|
||||
if err := m.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !fc1.closed || !fc2.closed {
|
||||
t.Fatal("Close should close all underlying clients")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRequestIDContext(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
if got := RequestIDFromContext(ctx); got != "" {
|
||||
t.Fatalf("expected empty request id, got %q", got)
|
||||
}
|
||||
ctx = WithRequestID(ctx, "req-123")
|
||||
if got := RequestIDFromContext(ctx); got != "req-123" {
|
||||
t.Fatalf("expected req-123, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTransportAndAuthHeader(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"": "streamable-http",
|
||||
"http": "streamable-http",
|
||||
"streamable-http": "streamable-http",
|
||||
"SSE": "sse",
|
||||
"weird": "weird",
|
||||
}
|
||||
for in, want := range cases {
|
||||
if got := normalizeTransport(in); got != want {
|
||||
t.Errorf("normalizeTransport(%q) = %q, want %q", in, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
if got := authHeaderValue(""); got != "" {
|
||||
t.Fatalf("expected empty auth header, got %q", got)
|
||||
}
|
||||
if got := authHeaderValue("abc"); got != "Bearer abc" {
|
||||
t.Fatalf("expected Bearer prefix, got %q", got)
|
||||
}
|
||||
if got := authHeaderValue("Bearer abc"); got != "Bearer abc" {
|
||||
t.Fatalf("expected unchanged, got %q", got)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user