From b22ee74148bd0210d4b8289d38ae5bf0040c87da Mon Sep 17 00:00:00 2001 From: Lunny Xiao Date: Mon, 24 Aug 2026 00:14:34 -0700 Subject: [PATCH] feat(params): add structured argument binding helper Add params.Bind, which unmarshals a tool call's map[string]any args into a typed struct via JSON round-trip so JSON numbers land in the correct Go numeric field types, and enforces `required:"true"` struct tags with clear errors. Migrate the branch, tree, and file repo handlers to use it instead of repeated args["x"].(string)/!ok extraction, preserving existing validation behavior for each field. Co-Authored-By: Codet (GPT-5-Codex) --- operation/repo/branch.go | 61 +++++++++--------- operation/repo/file.go | 131 ++++++++++++++++++-------------------- operation/repo/tree.go | 26 ++++---- pkg/params/bind.go | 55 ++++++++++++++++ pkg/params/bind_test.go | 132 +++++++++++++++++++++++++++++++++++++++ 5 files changed, 288 insertions(+), 117 deletions(-) create mode 100644 pkg/params/bind.go create mode 100644 pkg/params/bind_test.go diff --git a/operation/repo/branch.go b/operation/repo/branch.go index 461752f..c6149d1 100644 --- a/operation/repo/branch.go +++ b/operation/repo/branch.go @@ -69,28 +69,26 @@ func init() { }) } +type createBranchArgs struct { + Owner string `json:"owner" required:"true"` + Repo string `json:"repo" required:"true"` + Branch string `json:"branch" required:"true"` + OldBranch string `json:"old_branch"` +} + func CreateBranchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { - owner, err := params.GetString(args, "owner") - if err != nil { + var in createBranchArgs + if err := params.Bind(args, &in); err != nil { return to.ErrorResult(err) } - repo, err := params.GetString(args, "repo") - if err != nil { - return to.ErrorResult(err) - } - branch, err := params.GetString(args, "branch") - if err != nil { - return to.ErrorResult(err) - } - oldBranch, _ := args["old_branch"].(string) client, err := gitea.ClientFromContext(ctx) if err != nil { return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) } - _, _, err = client.Repositories.CreateBranch(ctx, owner, repo, gitea_sdk.CreateBranchOption{ - BranchName: branch, - OldBranchName: oldBranch, + _, _, err = client.Repositories.CreateBranch(ctx, in.Owner, in.Repo, gitea_sdk.CreateBranchOption{ + BranchName: in.Branch, + OldBranchName: in.OldBranch, }) if err != nil { return to.ErrorResult(fmt.Errorf("create branch error: %v", err)) @@ -99,24 +97,22 @@ func CreateBranchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResu return to.TextResult("Branch Created") } +type deleteBranchArgs struct { + Owner string `json:"owner" required:"true"` + Repo string `json:"repo" required:"true"` + Branch string `json:"branch" required:"true"` +} + func DeleteBranchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { - owner, err := params.GetString(args, "owner") - if err != nil { - return to.ErrorResult(err) - } - repo, err := params.GetString(args, "repo") - if err != nil { - return to.ErrorResult(err) - } - branch, err := params.GetString(args, "branch") - if err != nil { + var in deleteBranchArgs + if err := params.Bind(args, &in); err != nil { return to.ErrorResult(err) } client, err := gitea.ClientFromContext(ctx) if err != nil { return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) } - _, _, err = client.Repositories.DeleteRepoBranch(ctx, owner, repo, branch) + _, _, err = client.Repositories.DeleteRepoBranch(ctx, in.Owner, in.Repo, in.Branch) if err != nil { return to.ErrorResult(fmt.Errorf("delete branch error: %v", err)) } @@ -124,13 +120,14 @@ func DeleteBranchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResu return to.TextResult("Branch Deleted") } +type listBranchesArgs struct { + Owner string `json:"owner" required:"true"` + Repo string `json:"repo" required:"true"` +} + func ListBranchesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { - owner, err := params.GetString(args, "owner") - if err != nil { - return to.ErrorResult(err) - } - repo, err := params.GetString(args, "repo") - if err != nil { + var in listBranchesArgs + if err := params.Bind(args, &in); err != nil { return to.ErrorResult(err) } page, pageSize := params.GetPagination(args, 30) @@ -144,7 +141,7 @@ func ListBranchesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResu if err != nil { return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) } - branches, _, err := client.Repositories.ListRepoBranches(ctx, owner, repo, opt) + branches, _, err := client.Repositories.ListRepoBranches(ctx, in.Owner, in.Repo, opt) if err != nil { return to.ErrorResult(fmt.Errorf("list branches error: %v", err)) } diff --git a/operation/repo/file.go b/operation/repo/file.go index af10c07..fd5402d 100644 --- a/operation/repo/file.go +++ b/operation/repo/file.go @@ -102,30 +102,28 @@ type ContentLine struct { Content string `json:"content"` } +type getFileContentArgs struct { + Owner string `json:"owner" required:"true"` + Repo string `json:"repo" required:"true"` + Ref string `json:"ref"` + Path string `json:"path" required:"true"` + WithLines bool `json:"withLines"` +} + func GetFileContentFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { - owner, err := params.GetString(args, "owner") - if err != nil { - return to.ErrorResult(err) - } - repo, err := params.GetString(args, "repo") - if err != nil { - return to.ErrorResult(err) - } - ref, _ := args["ref"].(string) - filePath, err := params.GetString(args, "path") - if err != nil { + var in getFileContentArgs + if err := params.Bind(args, &in); err != nil { return to.ErrorResult(err) } client, err := gitea.ClientFromContext(ctx) if err != nil { return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) } - content, _, err := client.Repositories.GetContents(ctx, owner, repo, ref, filePath) + content, _, err := client.Repositories.GetContents(ctx, in.Owner, in.Repo, in.Ref, in.Path) if err != nil { return to.ErrorResult(fmt.Errorf("get file err: %v", err)) } - withLines, _ := args["withLines"].(bool) - if withLines { + if in.WithLines { rawContent, err := base64.StdEncoding.DecodeString(*content.Content) if err != nil { return to.ErrorResult(fmt.Errorf("decode base64 content err: %v", err)) @@ -164,49 +162,45 @@ func GetFileContentFn(ctx context.Context, args map[string]any) (*mcp.CallToolRe return to.TextResult(slimContents(content)) } +type getDirContentArgs struct { + Owner string `json:"owner" required:"true"` + Repo string `json:"repo" required:"true"` + Ref string `json:"ref"` + Path string `json:"path" required:"true"` +} + func GetDirContentFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { - owner, err := params.GetString(args, "owner") - if err != nil { - return to.ErrorResult(err) - } - repo, err := params.GetString(args, "repo") - if err != nil { - return to.ErrorResult(err) - } - ref, _ := args["ref"].(string) - filePath, err := params.GetString(args, "path") - if err != nil { + var in getDirContentArgs + if err := params.Bind(args, &in); err != nil { return to.ErrorResult(err) } client, err := gitea.ClientFromContext(ctx) if err != nil { return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) } - content, _, err := client.Repositories.ListContents(ctx, owner, repo, ref, filePath) + content, _, err := client.Repositories.ListContents(ctx, in.Owner, in.Repo, in.Ref, in.Path) if err != nil { return to.ErrorResult(fmt.Errorf("get dir content err: %v", err)) } return to.TextResult(slimDirEntries(content)) } +type createOrUpdateFileArgs struct { + Owner string `json:"owner" required:"true"` + Repo string `json:"repo" required:"true"` + Path string `json:"path" required:"true"` + Content string `json:"content"` + Message string `json:"message"` + BranchName string `json:"branch_name"` + NewBranchName string `json:"new_branch_name"` + SHA string `json:"sha"` +} + func CreateOrUpdateFileFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { - owner, err := params.GetString(args, "owner") - if err != nil { + var in createOrUpdateFileArgs + if err := params.Bind(args, &in); err != nil { return to.ErrorResult(err) } - repo, err := params.GetString(args, "repo") - if err != nil { - return to.ErrorResult(err) - } - filePath, err := params.GetString(args, "path") - if err != nil { - return to.ErrorResult(err) - } - content, _ := args["content"].(string) - message, _ := args["message"].(string) - branchName, _ := args["branch_name"].(string) - newBranchName, _ := args["new_branch_name"].(string) - sha, _ := args["sha"].(string) client, err := gitea.ClientFromContext(ctx) if err != nil { @@ -214,20 +208,20 @@ func CreateOrUpdateFileFn(ctx context.Context, args map[string]any) (*mcp.CallTo } fileOpt := gitea_sdk.FileOptions{ - Message: message, - BranchName: branchName, - NewBranchName: newBranchName, + Message: in.Message, + BranchName: in.BranchName, + NewBranchName: in.NewBranchName, } - targetBranch := cmp.Or(newBranchName, branchName) + targetBranch := cmp.Or(in.NewBranchName, in.BranchName) - if sha != "" { + if in.SHA != "" { // Update existing file opt := gitea_sdk.UpdateFileOptions{ - SHA: sha, - Content: base64.StdEncoding.EncodeToString([]byte(content)), + SHA: in.SHA, + Content: base64.StdEncoding.EncodeToString([]byte(in.Content)), FileOptions: fileOpt, } - _, _, err = client.Repositories.UpdateFile(ctx, owner, repo, filePath, opt) + _, _, err = client.Repositories.UpdateFile(ctx, in.Owner, in.Repo, in.Path, opt) if err != nil { return to.ErrorResult(fmt.Errorf("update file err: %v", err)) } @@ -236,47 +230,42 @@ func CreateOrUpdateFileFn(ctx context.Context, args map[string]any) (*mcp.CallTo // Create new file opt := gitea_sdk.CreateFileOptions{ - Content: base64.StdEncoding.EncodeToString([]byte(content)), + Content: base64.StdEncoding.EncodeToString([]byte(in.Content)), FileOptions: fileOpt, } - _, _, err = client.Repositories.CreateFile(ctx, owner, repo, filePath, opt) + _, _, err = client.Repositories.CreateFile(ctx, in.Owner, in.Repo, in.Path, opt) if err != nil { return to.ErrorResult(fmt.Errorf("create file err: %v", err)) } return to.TextResult("Create file success on branch " + targetBranch) } +type deleteFileArgs struct { + Owner string `json:"owner" required:"true"` + Repo string `json:"repo" required:"true"` + Path string `json:"path" required:"true"` + Message string `json:"message"` + BranchName string `json:"branch_name"` + SHA string `json:"sha" required:"true"` +} + func DeleteFileFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { - owner, err := params.GetString(args, "owner") - if err != nil { - return to.ErrorResult(err) - } - repo, err := params.GetString(args, "repo") - if err != nil { - return to.ErrorResult(err) - } - filePath, err := params.GetString(args, "path") - if err != nil { - return to.ErrorResult(err) - } - message, _ := args["message"].(string) - branchName, _ := args["branch_name"].(string) - sha, err := params.GetString(args, "sha") - if err != nil { + var in deleteFileArgs + if err := params.Bind(args, &in); err != nil { return to.ErrorResult(err) } opt := gitea_sdk.DeleteFileOptions{ FileOptions: gitea_sdk.FileOptions{ - Message: message, - BranchName: branchName, + Message: in.Message, + BranchName: in.BranchName, }, - SHA: sha, + SHA: in.SHA, } client, err := gitea.ClientFromContext(ctx) if err != nil { return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) } - _, err = client.Repositories.DeleteFile(ctx, owner, repo, filePath, opt) + _, err = client.Repositories.DeleteFile(ctx, in.Owner, in.Repo, in.Path, opt) if err != nil { return to.ErrorResult(fmt.Errorf("delete file err: %v", err)) } diff --git a/operation/repo/tree.go b/operation/repo/tree.go index c10bcbd..d6c023c 100644 --- a/operation/repo/tree.go +++ b/operation/repo/tree.go @@ -37,20 +37,18 @@ func init() { }) } +type getRepoTreeArgs struct { + Owner string `json:"owner" required:"true"` + Repo string `json:"repo" required:"true"` + TreeSHA string `json:"tree_sha" required:"true"` + Recursive bool `json:"recursive"` +} + func GetRepoTreeFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { - owner, err := params.GetString(args, "owner") - if err != nil { + var in getRepoTreeArgs + if err := params.Bind(args, &in); err != nil { return to.ErrorResult(err) } - repo, err := params.GetString(args, "repo") - if err != nil { - return to.ErrorResult(err) - } - treeSHA, err := params.GetString(args, "tree_sha") - if err != nil { - return to.ErrorResult(err) - } - recursive, _ := args["recursive"].(bool) page, pageSize := params.GetPagination(args, 30) opt := gitea_sdk.ListTreeOptions{ @@ -58,14 +56,14 @@ func GetRepoTreeFn(ctx context.Context, args map[string]any) (*mcp.CallToolResul Page: page, PageSize: pageSize, }, - Ref: treeSHA, - Recursive: recursive, + Ref: in.TreeSHA, + Recursive: in.Recursive, } client, err := gitea.ClientFromContext(ctx) if err != nil { return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) } - tree, _, err := client.Git.GetTrees(ctx, owner, repo, opt) + tree, _, err := client.Git.GetTrees(ctx, in.Owner, in.Repo, opt) if err != nil { return to.ErrorResult(fmt.Errorf("get repository tree err: %v", err)) } diff --git a/pkg/params/bind.go b/pkg/params/bind.go new file mode 100644 index 0000000..a51f177 --- /dev/null +++ b/pkg/params/bind.go @@ -0,0 +1,55 @@ +package params + +import ( + "encoding/json" + "fmt" + "reflect" + "strings" +) + +// Bind decodes args into out, a pointer to a struct, replacing the repeated +// args["x"].(string)/!ok extraction pattern. It round-trips args through JSON +// so JSON numbers land in the correct Go numeric field types. +// +// Struct fields tagged `required:"true"` must be present in args; a missing +// key, or an empty string value for a string field, returns an error naming +// the field's json tag. +func Bind(args map[string]any, out any) error { + v := reflect.ValueOf(out) + if v.Kind() != reflect.Pointer || v.Elem().Kind() != reflect.Struct { + return fmt.Errorf("params.Bind: out must be a pointer to a struct, got %T", out) + } + + if err := checkRequiredFields(args, v.Elem().Type()); err != nil { + return err + } + + data, err := json.Marshal(args) + if err != nil { + return fmt.Errorf("params.Bind: marshal args: %w", err) + } + if err := json.Unmarshal(data, out); err != nil { + return fmt.Errorf("params.Bind: %w", err) + } + return nil +} + +func checkRequiredFields(args map[string]any, t reflect.Type) error { + for field := range t.Fields() { + if field.Tag.Get("required") != "true" { + continue + } + name, _, _ := strings.Cut(field.Tag.Get("json"), ",") + if name == "" || name == "-" { + continue + } + val, ok := args[name] + if !ok { + return fmt.Errorf("%s is required", name) + } + if s, isString := val.(string); isString && s == "" { + return fmt.Errorf("%s is required", name) + } + } + return nil +} diff --git a/pkg/params/bind_test.go b/pkg/params/bind_test.go new file mode 100644 index 0000000..89e58d9 --- /dev/null +++ b/pkg/params/bind_test.go @@ -0,0 +1,132 @@ +package params + +import ( + "strings" + "testing" +) + +func TestBind_RequiredFieldsPresent(t *testing.T) { + type args struct { + Owner string `json:"owner" required:"true"` + Repo string `json:"repo" required:"true"` + } + var out args + err := Bind(map[string]any{"owner": "gitea", "repo": "gitea-mcp"}, &out) + if err != nil { + t.Fatalf("Bind() unexpected error = %v", err) + } + if out.Owner != "gitea" || out.Repo != "gitea-mcp" { + t.Errorf("Bind() = %+v, want Owner=gitea Repo=gitea-mcp", out) + } +} + +func TestBind_RequiredFieldMissing(t *testing.T) { + type args struct { + Owner string `json:"owner" required:"true"` + } + var out args + err := Bind(map[string]any{}, &out) + if err == nil { + t.Fatal("Bind() expected error, got nil") + } + if !strings.Contains(err.Error(), "owner") { + t.Errorf("Bind() error = %v, want mentioning %q", err, "owner") + } +} + +func TestBind_RequiredStringFieldEmpty(t *testing.T) { + type args struct { + Owner string `json:"owner" required:"true"` + } + var out args + err := Bind(map[string]any{"owner": ""}, &out) + if err == nil { + t.Fatal("Bind() expected error for empty required string, got nil") + } +} + +func TestBind_OptionalFieldDefaultsToZeroValue(t *testing.T) { + type args struct { + Owner string `json:"owner" required:"true"` + OldBranch string `json:"old_branch"` + } + var out args + err := Bind(map[string]any{"owner": "gitea"}, &out) + if err != nil { + t.Fatalf("Bind() unexpected error = %v", err) + } + if out.OldBranch != "" { + t.Errorf("Bind() OldBranch = %q, want empty", out.OldBranch) + } +} + +func TestBind_NumericConversion(t *testing.T) { + type args struct { + Page int `json:"page"` + PerPage int64 `json:"per_page"` + } + var out args + err := Bind(map[string]any{"page": float64(2), "per_page": float64(40)}, &out) + if err != nil { + t.Fatalf("Bind() unexpected error = %v", err) + } + if out.Page != 2 || out.PerPage != 40 { + t.Errorf("Bind() = %+v, want Page=2 PerPage=40", out) + } +} + +func TestBind_Boolean(t *testing.T) { + type args struct { + Recursive bool `json:"recursive"` + } + var out args + err := Bind(map[string]any{"recursive": true}, &out) + if err != nil { + t.Fatalf("Bind() unexpected error = %v", err) + } + if !out.Recursive { + t.Errorf("Bind() Recursive = false, want true") + } +} + +func TestBind_Array(t *testing.T) { + type args struct { + Labels []string `json:"labels"` + IDs []int64 `json:"ids"` + } + var out args + err := Bind(map[string]any{ + "labels": []any{"bug", "help wanted"}, + "ids": []any{float64(1), float64(2)}, + }, &out) + if err != nil { + t.Fatalf("Bind() unexpected error = %v", err) + } + if len(out.Labels) != 2 || out.Labels[0] != "bug" || out.Labels[1] != "help wanted" { + t.Errorf("Bind() Labels = %v, want [bug help wanted]", out.Labels) + } + if len(out.IDs) != 2 || out.IDs[0] != 1 || out.IDs[1] != 2 { + t.Errorf("Bind() IDs = %v, want [1 2]", out.IDs) + } +} + +func TestBind_InvalidFieldType(t *testing.T) { + type args struct { + Page int `json:"page"` + } + var out args + err := Bind(map[string]any{"page": "not-a-number"}, &out) + if err == nil { + t.Fatal("Bind() expected error for invalid numeric field, got nil") + } +} + +func TestBind_NonPointerRejected(t *testing.T) { + type args struct { + Owner string `json:"owner"` + } + err := Bind(map[string]any{}, args{}) + if err == nil { + t.Fatal("Bind() expected error for non-pointer out, got nil") + } +}