package toolkit import ( "testing" "github.com/modelcontextprotocol/go-sdk/mcp" ) func TestRequestArgs(t *testing.T) { req := &mcp.CallToolRequest{Params: &mcp.CallToolParamsRaw{ Arguments: []byte(`{"path":"/tmp/x","n":3}`), }} got := RequestArgs(req) if got["path"] != "/tmp/x" { t.Errorf("path = %v", got["path"]) } if RequestArgs(nil)["missing"] != nil { t.Error("nil request must yield empty map") } bad := &mcp.CallToolRequest{Params: &mcp.CallToolParamsRaw{Arguments: []byte(`{not json`)}} if len(RequestArgs(bad)) != 0 { t.Error("bad json must yield empty map") } } func TestScalarGetters(t *testing.T) { args := map[string]any{ "s": "hello", "n": float64(42), "ns": "17", "b": true, "null": nil, } if got := GetString(args, "s", "def"); got != "hello" { t.Errorf("GetString = %q", got) } if got := GetString(args, "n", "def"); got != "42" { t.Errorf("GetString(number) = %q", got) } if got := GetString(args, "missing", "def"); got != "def" { t.Errorf("GetString(missing) = %q", got) } if got := GetInt(args, "n", 0); got != 42 { t.Errorf("GetInt(float64) = %d", got) } if got := GetInt(args, "ns", 0); got != 17 { t.Errorf("GetInt(string) = %d", got) } if got := GetInt(args, "missing", 9); got != 9 { t.Errorf("GetInt(default) = %d", got) } if !GetBool(args, "b", false) { t.Error("GetBool = false, want true") } if GetBool(args, "missing", true) != true { t.Error("GetBool(default) = false") } } func TestRequireString(t *testing.T) { if _, err := RequireString(map[string]any{}, "x"); err == nil { t.Error("missing arg must error") } if _, err := RequireString(map[string]any{"x": ""}, "x"); err == nil { t.Error("empty arg must error") } if v, err := RequireString(map[string]any{"x": "ok"}, "x"); err != nil || v != "ok" { t.Errorf("RequireString = (%q, %v)", v, err) } } func TestArrayGetters(t *testing.T) { args := map[string]any{ "strs": []any{"a", "", "b", 3}, "ints": []any{float64(1), 2, "x"}, } strs, ok := GetStringArray(args, "strs") if !ok || len(strs) != 2 || strs[0] != "a" || strs[1] != "b" { t.Errorf("GetStringArray = (%v, %v)", strs, ok) } ints, ok := GetIntArray(args, "ints") if !ok || len(ints) != 2 || ints[0] != 1 || ints[1] != 2 { t.Errorf("GetIntArray = (%v, %v)", ints, ok) } if _, ok := GetStringArray(args, "missing"); ok { t.Error("missing array must be ok=false") } if _, ok := GetStringArray(map[string]any{"e": []any{}}, "e"); ok { t.Error("empty array must be ok=false") } } func TestGetStringMap(t *testing.T) { args := map[string]any{"h": map[string]any{"A": "1", "B": 2}} got := GetStringMap(args, "h") if got["A"] != "1" { t.Errorf("GetStringMap = %v", got) } if _, exists := got["B"]; exists { t.Error("non-string value must be dropped") } if GetStringMap(args, "missing") != nil { t.Error("missing map must be nil") } }