132 lines
3.8 KiB
Go
132 lines
3.8 KiB
Go
package mcpc
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"testing"
|
|
|
|
"github.com/modelcontextprotocol/go-sdk/mcp"
|
|
|
|
"nub/internal/config"
|
|
"nub/internal/tool"
|
|
)
|
|
|
|
type echoArgs struct {
|
|
Text string `json:"text"`
|
|
}
|
|
|
|
// startTestServer läuft komplett in-process über mcp.NewInMemoryTransports —
|
|
// ein echter MCP-Server/-Client-Roundtrip ohne Subprozess oder Netzwerk.
|
|
func startTestServer(t *testing.T, toolNames ...string) *mcp.ClientSession {
|
|
t.Helper()
|
|
if len(toolNames) == 0 {
|
|
toolNames = []string{"echo"}
|
|
}
|
|
|
|
server := mcp.NewServer(&mcp.Implementation{Name: "test-server"}, nil)
|
|
for _, name := range toolNames {
|
|
name := name
|
|
mcp.AddTool(server, &mcp.Tool{Name: name, Description: "echoes text"},
|
|
func(ctx context.Context, req *mcp.CallToolRequest, args echoArgs) (*mcp.CallToolResult, any, error) {
|
|
return &mcp.CallToolResult{Content: []mcp.Content{&mcp.TextContent{Text: name + ":" + args.Text}}}, nil, nil
|
|
})
|
|
}
|
|
|
|
serverTransport, clientTransport := mcp.NewInMemoryTransports()
|
|
ctx := context.Background()
|
|
go func() { _ = server.Run(ctx, serverTransport) }()
|
|
|
|
client := mcp.NewClient(&mcp.Implementation{Name: "test-client"}, nil)
|
|
session, err := client.Connect(ctx, clientTransport, nil)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
t.Cleanup(func() { _ = session.Close() })
|
|
return session
|
|
}
|
|
|
|
func TestBridge_RegisterAndCallTool(t *testing.T) {
|
|
session := startTestServer(t, "echo")
|
|
srv := &Server{Name: "testsrv", Session: session}
|
|
|
|
reg := tool.NewRegistry()
|
|
warnings, err := RegisterTools(context.Background(), reg, srv, config.MCPServer{Name: "testsrv"})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(warnings) != 0 {
|
|
t.Errorf("unexpected warnings: %v", warnings)
|
|
}
|
|
|
|
got, ok := reg.Get("testsrv__echo")
|
|
if !ok {
|
|
t.Fatal("expected tool testsrv__echo to be registered with server-prefixed name")
|
|
}
|
|
|
|
input, _ := json.Marshal(map[string]string{"text": "hi"})
|
|
res, err := got.Run(context.Background(), input, tool.Env{})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if res.IsError {
|
|
t.Fatalf("unexpected error result: %s", res.ForModel)
|
|
}
|
|
if res.ForModel != "echo:hi" {
|
|
t.Errorf("ForModel = %q, want echo:hi", res.ForModel)
|
|
}
|
|
}
|
|
|
|
func TestBridge_AllowlistFiltersTools(t *testing.T) {
|
|
session := startTestServer(t, "echo", "danger")
|
|
srv := &Server{Name: "testsrv", Session: session}
|
|
|
|
reg := tool.NewRegistry()
|
|
_, err := RegisterTools(context.Background(), reg, srv, config.MCPServer{Name: "testsrv", Tools: []string{"echo"}})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
if _, ok := reg.Get("testsrv__echo"); !ok {
|
|
t.Error("expected allowed tool to be registered")
|
|
}
|
|
if _, ok := reg.Get("testsrv__danger"); ok {
|
|
t.Error("expected non-allowlisted tool to be filtered out")
|
|
}
|
|
}
|
|
|
|
func TestBridge_RuntimeFailureIsErrorNotPanic(t *testing.T) {
|
|
session := startTestServer(t, "echo")
|
|
srv := &Server{Name: "testsrv", Session: session}
|
|
|
|
reg := tool.NewRegistry()
|
|
if _, err := RegisterTools(context.Background(), reg, srv, config.MCPServer{Name: "testsrv"}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
got, _ := reg.Get("testsrv__echo")
|
|
|
|
// Server-Ausfall simulieren: Session vor dem Aufruf schließen.
|
|
if err := session.Close(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
input, _ := json.Marshal(map[string]string{"text": "hi"})
|
|
res, err := got.Run(context.Background(), input, tool.Env{})
|
|
if err != nil {
|
|
t.Fatalf("Run should not return a Go error on server failure, got: %v", err)
|
|
}
|
|
if !res.IsError {
|
|
t.Error("expected IsError=true when the underlying MCP session is closed")
|
|
}
|
|
}
|
|
|
|
func TestConnectAll_BrokenServerDoesNotBlockStart(t *testing.T) {
|
|
result := ConnectAll(context.Background(), []config.MCPServer{
|
|
{Name: "broken", Command: "this-binary-does-not-exist-xyz"},
|
|
})
|
|
if len(result.Servers) != 0 {
|
|
t.Errorf("expected no connected servers, got %d", len(result.Servers))
|
|
}
|
|
if len(result.Warnings) != 1 {
|
|
t.Fatalf("expected exactly one warning, got %v", result.Warnings)
|
|
}
|
|
}
|