Compare commits

..

5 Commits

Author SHA1 Message Date
Bo-Yi Wu (吳柏毅) 7de5f1021f Merge branch 'main' into refactor/migrate-official-mcp-go-sdk 2026-08-03 04:17:20 +00:00
Bo-Yi Wu cc0cb109a8 fix(mcp): address SDK migration review
Co-Authored-By: OpenAI Codex (GPT-5) <noreply@openai.com>
2026-08-03 12:02:21 +08:00
silverwind 4eaeb252a0 docs: sync AGENTS.md with gitea/gitea (#224)
Carries over the applicable rules from https://gitea.com/gitea/gitea `AGENTS.md`, covering test scope, linter escape hatches and PR conventions. Swaps `Co-Authored-By` for the `Assisted-by` trailer.

Reviewed-on: https://gitea.com/gitea/gitea-mcp/pulls/224
Reviewed-by: Lunny Xiao <xiaolunwen@gmail.com>
Co-authored-by: silverwind <me@silverwind.io>
2026-08-02 20:28:58 +00:00
silverwind 0dc9868e2e fix: accept null tool arguments and bound HTTP resource use
Review follow-ups on the SDK migration.

An "arguments": null is what clients send for parameterless tools like get_me,
and what mcp-go accepted by returning a nil map. The new adapter rejected it
with InvalidParams, which broke those calls outright.

The /mcp endpoint took unlimited request bodies and never expired idle
sessions, so a peer that goes away without DELETE kept its session for the
process lifetime. Both are reachable before any token check, so neither can
stay unbounded; the body cap sits above the SDK default to leave room for the
base64 content create_or_update_file accepts.

Required() smuggled a bool through the property schema map and deleted it
again, colliding with the JSON Schema keyword of the same name. It now sets a
field on Property, so an object property can carry its own required list.

The tool contract fixture cost a manual regeneration step and four
hand-maintained counts on every tool change, and a snapshot freezes defects
rather than reporting them. Property assertions cover the same surface and
reject a duplicate tool name, a readOnlyHint that disagrees with the register
call, and a default that contradicts its own type or enum.

Co-Authored-By: Claude (Opus 5) <noreply@anthropic.com>
2026-08-02 19:34:40 +02:00
Bo-Yi Wu 80c8b25d6e refactor!: replace mcp-go with the official MCP Go SDK
- Swap the mark3labs MCP dependency for modelcontextprotocol/go-sdk v1.7.0
- Add a declarative tool definition and JSON Schema builder to the tool package
- Adapt registered tools to the official low-level handler, recovering panics and mapping failures to JSON-RPC errors
- Narrow tool handlers to take a plain argument map instead of an SDK request type
- Rewire stdio and HTTP transports onto the official server, moving Authorization parsing into receiving middleware
- Add a golden contract test that locks the exposed tool schemas, plus SDK integration and helper tests
- Add a test target and run it in the pull request workflow

BREAKING CHANGE: The exported helpers change signature. Tool.RegisterRead and
Tool.RegisterWrite now take tool.ServerTool instead of server.ServerTool, and the
annotation constructors return *mcp.ToolAnnotations from the official SDK. Callers
must build tool definitions with tool.NewDefinition and handlers with the
func(context.Context, map[string]any) signature.

The 30-second SSE heartbeat is removed because the official SDK has no equivalent.
ServerOptions.KeepAlive is deliberately not used as a substitute, since it sends MCP
ping requests and ping is removed in protocol 2026. HTTP stays stateless=false, so
new clients negotiate at most 2025-11-25.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-08-02 21:30:19 +08:00
49 changed files with 2477 additions and 1578 deletions
+2
View File
@@ -13,6 +13,8 @@ jobs:
go-version-file: 'go.mod' go-version-file: 'go.mod'
- name: lint - name: lint
run: make lint run: make lint
- name: test
run: make test
- name: build - name: build
run: make build run: make build
- name: security-check - name: security-check
+12 -8
View File
@@ -1,12 +1,16 @@
- Never assume, verify before claiming
- Use `make help` to find available development targets - Use `make help` to find available development targets
- Run `make fmt` to format `.go` files, and run `make lint-go` to lint them - PR descriptions: minimal, only what and why, no task lists or file listings
- Run `make tidy` after any `go.mod` changes - Reference issues and PRs by full URL, not by number
- Run single go tests with `go test -run '^TestName$' ./modulepath/`
- Ensure no trailing whitespace in edited files
- Use Conventional Commits for commit messages and PR titles, e.g. `type(scope): subject`; `!` before the colon if breaking. Use `test` type for test-only changes. - Use Conventional Commits for commit messages and PR titles, e.g. `type(scope): subject`; `!` before the colon if breaking. Use `test` type for test-only changes.
- Add an `Assisted-by: AGENT_NAME:MODEL_VERSION` trailer to commit messages, never `Co-Authored-By` or `Signed-off-by`
- Attribute agent authorship on one trailing line in issue and pull request comments, never as a PR description section
- Never force-push, amend, or squash unless asked. Use new commits and normal push for pull request updates - Never force-push, amend, or squash unless asked. Use new commits and normal push for pull request updates
- Preserve existing code comments, do not remove or rewrite comments that are still relevant - Keep comments short, prefer same-line, explain why, never narrate code. Preserve existing ones that still apply
- Keep comments short, prefer same-line, explain why, never narrate code - Ensure no trailing whitespace in edited files
- Run `make fmt` to format `.go` files, `make lint-go` to lint them, and `make tidy` after any `go.mod` changes
- Fix the cause rather than disabling a linter or weakening a test. Where unavoidable, use the narrowest scope with a trailing comment giving the reason
- Register new tools with `Tool.RegisterRead` or `Tool.RegisterWrite`, and add them to the tool tables in `README.md`, `README.zh-cn.md` and `README.zh-tw.md` - Register new tools with `Tool.RegisterRead` or `Tool.RegisterWrite`, and add them to the tool tables in `README.md`, `README.zh-cn.md` and `README.zh-tw.md`
- Include authorship attribution in issue and pull request comments - Run single go tests with `go test -run '^TestName$' ./modulepath/`
- Add `Co-Authored-By` lines to all commits, indicating name and model used - Write the fewest, fastest tests covering the behavior, extending an existing one where possible. Prefer unit tests where logic is testable in isolation
- Wait on a deterministic condition rather than `sleep`
+6
View File
@@ -6,6 +6,8 @@ LDFLAGS := -X "main.Version=$(VERSION)"
GOLANGCI_LINT_PACKAGE ?= github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.12.2 # renovate: datasource=go GOLANGCI_LINT_PACKAGE ?= github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.12.2 # renovate: datasource=go
GOVULNCHECK_PACKAGE ?= golang.org/x/vuln/cmd/govulncheck@v1.3.0 # renovate: datasource=go GOVULNCHECK_PACKAGE ?= golang.org/x/vuln/cmd/govulncheck@v1.3.0 # renovate: datasource=go
GOTEST_FLAGS ?= -race -timeout 20m
.PHONY: help .PHONY: help
help: ## print this help message help: ## print this help message
@echo "Usage: make [target]" @echo "Usage: make [target]"
@@ -38,6 +40,10 @@ clean: ## delete build artifacts
build: ## build the application build: ## build the application
$(GO) build -v -ldflags '-s -w $(LDFLAGS)' -o $(EXECUTABLE) $(GO) build -v -ldflags '-s -w $(LDFLAGS)' -o $(EXECUTABLE)
.PHONY: test
test: ## run Go tests
$(GO) test $(GOTEST_FLAGS) ./...
.PHONY: air .PHONY: air
air: ## install air for hot reload air: ## install air for hot reload
@hash air > /dev/null 2>&1; if [ $$? -ne 0 ]; then \ @hash air > /dev/null 2>&1; if [ $$? -ne 0 ]; then \
+7 -5
View File
@@ -4,7 +4,7 @@ go 1.26.0
require ( require (
gitea.dev/sdk v1.2.0 gitea.dev/sdk v1.2.0
github.com/mark3labs/mcp-go v0.56.0 github.com/modelcontextprotocol/go-sdk v1.7.0
go.uber.org/zap v1.28.0 go.uber.org/zap v1.28.0
go.uber.org/zap/exp v0.3.0 go.uber.org/zap/exp v0.3.0
gopkg.in/natefinch/lumberjack.v2 v2.2.1 gopkg.in/natefinch/lumberjack.v2 v2.2.1
@@ -14,13 +14,15 @@ require (
github.com/42wim/httpsig v1.2.4 // indirect github.com/42wim/httpsig v1.2.4 // indirect
github.com/davidmz/go-pageant v1.0.2 // indirect github.com/davidmz/go-pageant v1.0.2 // indirect
github.com/google/jsonschema-go v0.4.3 // indirect github.com/google/jsonschema-go v0.4.3 // indirect
github.com/google/uuid v1.6.0 // indirect
github.com/hashicorp/go-version v1.9.0 // indirect github.com/hashicorp/go-version v1.9.0 // indirect
github.com/santhosh-tekuri/jsonschema/v6 v6.0.2 // indirect github.com/segmentio/asm v1.1.3 // indirect
github.com/spf13/cast v1.10.0 // indirect github.com/segmentio/encoding v0.5.4 // indirect
github.com/yosida95/uritemplate/v3 v3.0.2 // indirect github.com/yosida95/uritemplate/v3 v3.0.2 // indirect
go.uber.org/multierr v1.11.0 // indirect go.uber.org/multierr v1.11.0 // indirect
golang.org/x/crypto v0.54.0 // indirect golang.org/x/crypto v0.54.0 // indirect
golang.org/x/oauth2 v0.35.0 // indirect
golang.org/x/sync v0.22.0 // indirect
golang.org/x/sys v0.47.0 // indirect golang.org/x/sys v0.47.0 // indirect
golang.org/x/text v0.40.0 // indirect golang.org/x/time v0.15.0 // indirect
golang.org/x/tools v0.47.0 // indirect
) )
+16 -20
View File
@@ -6,32 +6,22 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davidmz/go-pageant v1.0.2 h1:bPblRCh5jGU+Uptpz6LgMZGD5hJoOt7otgT454WvHn0= github.com/davidmz/go-pageant v1.0.2 h1:bPblRCh5jGU+Uptpz6LgMZGD5hJoOt7otgT454WvHn0=
github.com/davidmz/go-pageant v1.0.2/go.mod h1:P2EDDnMqIwG5Rrp05dTRITj9z2zpGcD9efWSkTNKLIE= github.com/davidmz/go-pageant v1.0.2/go.mod h1:P2EDDnMqIwG5Rrp05dTRITj9z2zpGcD9efWSkTNKLIE=
github.com/dlclark/regexp2 v1.11.0 h1:G/nrcoOa7ZXlpoa/91N3X7mM3r8eIlMBBJZvsz/mxKI= github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
github.com/dlclark/regexp2 v1.11.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8= github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0= github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0=
github.com/google/jsonschema-go v0.4.3/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE= github.com/google/jsonschema-go v0.4.3/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/hashicorp/go-version v1.9.0 h1:CeOIz6k+LoN3qX9Z0tyQrPtiB1DFYRPfCIBtaXPSCnA= github.com/hashicorp/go-version v1.9.0 h1:CeOIz6k+LoN3qX9Z0tyQrPtiB1DFYRPfCIBtaXPSCnA=
github.com/hashicorp/go-version v1.9.0/go.mod h1:fltr4n8CU8Ke44wwGCBoEymUuxUHl09ZGVZPK5anwXA= github.com/hashicorp/go-version v1.9.0/go.mod h1:fltr4n8CU8Ke44wwGCBoEymUuxUHl09ZGVZPK5anwXA=
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/modelcontextprotocol/go-sdk v1.7.0 h1:yqjY2dsbKAC0LSuWZVBMrHgiG8ukXv6NRo0JiALay44=
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= github.com/modelcontextprotocol/go-sdk v1.7.0/go.mod h1:dL7u98E/zjJTGzEq+j30jQ8K2k1mb6LeAH4inEcSGts=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
github.com/mark3labs/mcp-go v0.56.0 h1:7aCj2wODCskMi08f923ADG+EfELZBdiKILny415cIS8=
github.com/mark3labs/mcp-go v0.56.0/go.mod h1:+8WclSK1ZUweCP3hvktSji8n8ABG/95QaEkeVE/Uwas=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= github.com/segmentio/asm v1.1.3 h1:WM03sfUOENvvKexOLp+pCqgb/WDjsi7EK8gIsICtzhc=
github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/segmentio/asm v1.1.3/go.mod h1:Ld3L4ZXGNcSLRg4JBsZ3//1+f/TjYl0Mzen/DQy1EJg=
github.com/santhosh-tekuri/jsonschema/v6 v6.0.2 h1:KRzFb2m7YtdldCEkzs6KqmJw4nqEVZGK7IN2kJkjTuQ= github.com/segmentio/encoding v0.5.4 h1:OW1VRern8Nw6ITAtwSZ7Idrl3MXCFwXHPgqESYfvNt0=
github.com/santhosh-tekuri/jsonschema/v6 v6.0.2/go.mod h1:JXeL+ps8p7/KNMjDQk3TCwPpBy0wYklyWTfbkIzdIFU= github.com/segmentio/encoding v0.5.4/go.mod h1:HS1ZKa3kSN32ZHVZ7ZLPLXWvOVIiZtyJnO1gPH1sKt0=
github.com/spf13/cast v1.10.0 h1:h2x0u2shc1QuLHfxi+cTJvs30+ZAHOGRic8uyGTDWxY=
github.com/spf13/cast v1.10.0/go.mod h1:jNfB8QC9IA6ZuY2ZjDp0KtFO2LZZlg4S/7bzP6qqeHo=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4= github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4=
@@ -50,6 +40,10 @@ golang.org/x/crypto v0.0.0-20210513164829-c07d793c2f9a/go.mod h1:P+XmwS30IXTQdn5
golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw= golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
golang.org/x/oauth2 v0.35.0 h1:Mv2mzuHuZuY2+bkyWXIHMfhNdJAdwW3FuWeCPYN5GVQ=
golang.org/x/oauth2 v0.35.0/go.mod h1:lzm5WQJQwKZ3nwavOZ3IS5Aulzxi68dUSgRHujetwEA=
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
@@ -57,9 +51,11 @@ golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9sn
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0= golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w= golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q=
golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA=
gopkg.in/natefinch/lumberjack.v2 v2.2.1 h1:bBRl1b0OH9s/DuPhuXpNl+VtCaJXFZ5/uEFST95x9zc= gopkg.in/natefinch/lumberjack.v2 v2.2.1 h1:bBRl1b0OH9s/DuPhuXpNl+VtCaJXFZ5/uEFST95x9zc=
gopkg.in/natefinch/lumberjack.v2 v2.2.1/go.mod h1:YD8tP3GAjkrDg1eZH7EGmyESg/lsYskCTPBJVb9jqSc= gopkg.in/natefinch/lumberjack.v2 v2.2.1/go.mod h1:YD8tP3GAjkrDg1eZH7EGmyESg/lsYskCTPBJVb9jqSc=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
+21 -21
View File
@@ -15,7 +15,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/params" "gitea.com/gitea/gitea-mcp/pkg/params"
"gitea.com/gitea/gitea-mcp/pkg/to" "gitea.com/gitea/gitea-mcp/pkg/to"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
// Artifact endpoints require Gitea 1.25+. Older servers answer 404/405, which is // Artifact endpoints require Gitea 1.25+. Older servers answer 404/405, which is
@@ -28,21 +28,21 @@ func artifactNotSupportedErr(err error) error {
return err return err
} }
func listRepoActionArtifactsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoActionArtifactsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
query := url.Values{} query := url.Values{}
query.Set("page", strconv.Itoa(page)) query.Set("page", strconv.Itoa(page))
query.Set("limit", strconv.Itoa(pageSize)) query.Set("limit", strconv.Itoa(pageSize))
if name := params.GetOptionalString(req.GetArguments(), "artifact_name", ""); name != "" { if name := params.GetOptionalString(args, "artifact_name", ""); name != "" {
query.Set("name", name) query.Set("name", name)
} }
@@ -59,25 +59,25 @@ func listRepoActionArtifactsFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(slimActionArtifacts(result)) return to.TextResult(slimActionArtifacts(result))
} }
func listRepoActionRunArtifactsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoActionRunArtifactsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
runID, err := params.GetIndex(req.GetArguments(), "run_id") runID, err := params.GetIndex(args, "run_id")
if err != nil || runID <= 0 { if err != nil || runID <= 0 {
return to.ErrorResult(errors.New("run_id is required")) return to.ErrorResult(errors.New("run_id is required"))
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
query := url.Values{} query := url.Values{}
query.Set("page", strconv.Itoa(page)) query.Set("page", strconv.Itoa(page))
query.Set("limit", strconv.Itoa(pageSize)) query.Set("limit", strconv.Itoa(pageSize))
if name := params.GetOptionalString(req.GetArguments(), "artifact_name", ""); name != "" { if name := params.GetOptionalString(args, "artifact_name", ""); name != "" {
query.Set("name", name) query.Set("name", name)
} }
@@ -94,16 +94,16 @@ func listRepoActionRunArtifactsFn(ctx context.Context, req mcp.CallToolRequest)
return to.TextResult(slimActionArtifacts(result)) return to.TextResult(slimActionArtifacts(result))
} }
func getRepoActionArtifactFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getRepoActionArtifactFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
artifactID, err := params.GetIndex(req.GetArguments(), "artifact_id") artifactID, err := params.GetIndex(args, "artifact_id")
if err != nil || artifactID <= 0 { if err != nil || artifactID <= 0 {
return to.ErrorResult(errors.New("artifact_id is required")) return to.ErrorResult(errors.New("artifact_id is required"))
} }
@@ -121,20 +121,20 @@ func getRepoActionArtifactFn(ctx context.Context, req mcp.CallToolRequest) (*mcp
return to.TextResult(slimActionArtifact(result)) return to.TextResult(slimActionArtifact(result))
} }
func downloadRepoActionArtifactFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func downloadRepoActionArtifactFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
artifactID, err := params.GetIndex(req.GetArguments(), "artifact_id") artifactID, err := params.GetIndex(args, "artifact_id")
if err != nil || artifactID <= 0 { if err != nil || artifactID <= 0 {
return to.ErrorResult(errors.New("artifact_id is required")) return to.ErrorResult(errors.New("artifact_id is required"))
} }
outputPath, _ := req.GetArguments()["output_path"].(string) outputPath, _ := args["output_path"].(string)
// Best-effort metadata lookup: gives a friendly filename and lets us fail // Best-effort metadata lookup: gives a friendly filename and lets us fail
// early with a clear message when the artifact has expired. // early with a clear message when the artifact has expired.
+111 -111
View File
@@ -11,10 +11,10 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/gitea" "gitea.com/gitea/gitea-mcp/pkg/gitea"
"gitea.com/gitea/gitea-mcp/pkg/params" "gitea.com/gitea/gitea-mcp/pkg/params"
"gitea.com/gitea/gitea-mcp/pkg/to" "gitea.com/gitea/gitea-mcp/pkg/to"
"gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
const ( const (
@@ -44,103 +44,103 @@ func toSecretMetas(secrets []*gitea_sdk.Secret) []secretMeta {
} }
var ( var (
ActionsConfigReadTool = mcp.NewTool( ActionsConfigReadTool = tool.NewDefinition(
ActionsConfigReadToolName, ActionsConfigReadToolName,
mcp.WithDescription("Read Actions secrets and variables."), "Read Actions secrets and variables.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read Actions secrets and variables")), annotation.ReadOnly("Read Actions secrets and variables"),
mcp.WithString("method", mcp.Required(), mcp.Enum("list_repo_secrets", "list_org_secrets", "list_repo_variables", "get_repo_variable", "list_org_variables", "get_org_variable")), tool.String("method", tool.Required(), tool.Enum("list_repo_secrets", "list_org_secrets", "list_repo_variables", "get_repo_variable", "list_org_variables", "get_org_variable")),
mcp.WithString("owner", mcp.Description("for repo methods")), tool.String("owner", tool.Description("for repo methods")),
mcp.WithString("repo", mcp.Description("for repo methods")), tool.String("repo", tool.Description("for repo methods")),
mcp.WithString("org", mcp.Description("for org methods")), tool.String("org", tool.Description("for org methods")),
mcp.WithString("name", mcp.Description("for get methods")), tool.String("name", tool.Description("for get methods")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1), mcp.Min(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1), tool.Minimum(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30), mcp.Min(1)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30), tool.Minimum(1)),
) )
ActionsConfigWriteTool = mcp.NewTool( ActionsConfigWriteTool = tool.NewDefinition(
ActionsConfigWriteToolName, ActionsConfigWriteToolName,
mcp.WithDescription("Write Actions secrets and variables: upsert, create, update, delete."), "Write Actions secrets and variables: upsert, create, update, delete.",
mcp.WithToolAnnotation(annotation.Destructive("Manage Actions secrets and variables")), annotation.Destructive("Manage Actions secrets and variables"),
mcp.WithString("method", mcp.Required(), mcp.Enum("upsert_repo_secret", "delete_repo_secret", "upsert_org_secret", "delete_org_secret", "create_repo_variable", "update_repo_variable", "delete_repo_variable", "create_org_variable", "update_org_variable", "delete_org_variable")), tool.String("method", tool.Required(), tool.Enum("upsert_repo_secret", "delete_repo_secret", "upsert_org_secret", "delete_org_secret", "create_repo_variable", "update_repo_variable", "delete_repo_variable", "create_org_variable", "update_org_variable", "delete_org_variable")),
mcp.WithString("owner", mcp.Description("for repo methods")), tool.String("owner", tool.Description("for repo methods")),
mcp.WithString("repo", mcp.Description("for repo methods")), tool.String("repo", tool.Description("for repo methods")),
mcp.WithString("org", mcp.Description("for org methods")), tool.String("org", tool.Description("for org methods")),
mcp.WithString("name", mcp.Description("secret or variable name")), tool.String("name", tool.Description("secret or variable name")),
mcp.WithString("data", mcp.Description("secret value (upsert)")), tool.String("data", tool.Description("secret value (upsert)")),
mcp.WithString("value", mcp.Description("variable value")), tool.String("value", tool.Description("variable value")),
mcp.WithString("description"), tool.String("description"),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{Tool: ActionsConfigReadTool, Handler: configReadFn}) Tool.RegisterRead(tool.ServerTool{Tool: ActionsConfigReadTool, Handler: configReadFn})
Tool.RegisterWrite(server.ServerTool{Tool: ActionsConfigWriteTool, Handler: configWriteFn}) Tool.RegisterWrite(tool.ServerTool{Tool: ActionsConfigWriteTool, Handler: configWriteFn})
} }
func configReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func configReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "list_repo_secrets": case "list_repo_secrets":
return listRepoActionSecretsFn(ctx, req) return listRepoActionSecretsFn(ctx, args)
case "list_org_secrets": case "list_org_secrets":
return listOrgActionSecretsFn(ctx, req) return listOrgActionSecretsFn(ctx, args)
case "list_repo_variables": case "list_repo_variables":
return listRepoActionVariablesFn(ctx, req) return listRepoActionVariablesFn(ctx, args)
case "get_repo_variable": case "get_repo_variable":
return getRepoActionVariableFn(ctx, req) return getRepoActionVariableFn(ctx, args)
case "list_org_variables": case "list_org_variables":
return listOrgActionVariablesFn(ctx, req) return listOrgActionVariablesFn(ctx, args)
case "get_org_variable": case "get_org_variable":
return getOrgActionVariableFn(ctx, req) return getOrgActionVariableFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func configWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func configWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "upsert_repo_secret": case "upsert_repo_secret":
return upsertRepoActionSecretFn(ctx, req) return upsertRepoActionSecretFn(ctx, args)
case "delete_repo_secret": case "delete_repo_secret":
return deleteRepoActionSecretFn(ctx, req) return deleteRepoActionSecretFn(ctx, args)
case "upsert_org_secret": case "upsert_org_secret":
return upsertOrgActionSecretFn(ctx, req) return upsertOrgActionSecretFn(ctx, args)
case "delete_org_secret": case "delete_org_secret":
return deleteOrgActionSecretFn(ctx, req) return deleteOrgActionSecretFn(ctx, args)
case "create_repo_variable": case "create_repo_variable":
return createRepoActionVariableFn(ctx, req) return createRepoActionVariableFn(ctx, args)
case "update_repo_variable": case "update_repo_variable":
return updateRepoActionVariableFn(ctx, req) return updateRepoActionVariableFn(ctx, args)
case "delete_repo_variable": case "delete_repo_variable":
return deleteRepoActionVariableFn(ctx, req) return deleteRepoActionVariableFn(ctx, args)
case "create_org_variable": case "create_org_variable":
return createOrgActionVariableFn(ctx, req) return createOrgActionVariableFn(ctx, args)
case "update_org_variable": case "update_org_variable":
return updateOrgActionVariableFn(ctx, req) return updateOrgActionVariableFn(ctx, args)
case "delete_org_variable": case "delete_org_variable":
return deleteOrgActionVariableFn(ctx, req) return deleteOrgActionVariableFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func listRepoActionSecretsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoActionSecretsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
@@ -157,24 +157,24 @@ func listRepoActionSecretsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp
return to.TextResult(toSecretMetas(secrets)) return to.TextResult(toSecretMetas(secrets))
} }
func upsertRepoActionSecretFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func upsertRepoActionSecretFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
data, err := params.GetString(req.GetArguments(), "data") data, err := params.GetString(args, "data")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
description, _ := req.GetArguments()["description"].(string) description, _ := args["description"].(string)
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
@@ -190,16 +190,16 @@ func upsertRepoActionSecretFn(ctx context.Context, req mcp.CallToolRequest) (*mc
return to.TextResult(map[string]any{"message": "secret upserted", "status": resp.StatusCode}) return to.TextResult(map[string]any{"message": "secret upserted", "status": resp.StatusCode})
} }
func deleteRepoActionSecretFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteRepoActionSecretFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -215,12 +215,12 @@ func deleteRepoActionSecretFn(ctx context.Context, req mcp.CallToolRequest) (*mc
return to.TextResult(map[string]any{"message": "secret deleted", "status": resp.StatusCode}) return to.TextResult(map[string]any{"message": "secret deleted", "status": resp.StatusCode})
} }
func listOrgActionSecretsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listOrgActionSecretsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
@@ -237,20 +237,20 @@ func listOrgActionSecretsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.
return to.TextResult(toSecretMetas(secrets)) return to.TextResult(toSecretMetas(secrets))
} }
func upsertOrgActionSecretFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func upsertOrgActionSecretFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
data, err := params.GetString(req.GetArguments(), "data") data, err := params.GetString(args, "data")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
description, _ := req.GetArguments()["description"].(string) description, _ := args["description"].(string)
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
@@ -266,12 +266,12 @@ func upsertOrgActionSecretFn(ctx context.Context, req mcp.CallToolRequest) (*mcp
return to.TextResult(map[string]any{"message": "secret upserted", "status": resp.StatusCode}) return to.TextResult(map[string]any{"message": "secret upserted", "status": resp.StatusCode})
} }
func deleteOrgActionSecretFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteOrgActionSecretFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -285,16 +285,16 @@ func deleteOrgActionSecretFn(ctx context.Context, req mcp.CallToolRequest) (*mcp
return to.TextResult(map[string]any{"message": "secret deleted"}) return to.TextResult(map[string]any{"message": "secret deleted"})
} }
func listRepoActionVariablesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoActionVariablesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
query := url.Values{} query := url.Values{}
query.Set("page", strconv.Itoa(page)) query.Set("page", strconv.Itoa(page))
@@ -308,16 +308,16 @@ func listRepoActionVariablesFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(result) return to.TextResult(result)
} }
func getRepoActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getRepoActionVariableFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -333,20 +333,20 @@ func getRepoActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp
return to.TextResult(variable) return to.TextResult(variable)
} }
func createRepoActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createRepoActionVariableFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
value, err := params.GetString(req.GetArguments(), "value") value, err := params.GetString(args, "value")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -362,20 +362,20 @@ func createRepoActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*
return to.TextResult(map[string]any{"message": "variable created", "status": resp.StatusCode}) return to.TextResult(map[string]any{"message": "variable created", "status": resp.StatusCode})
} }
func updateRepoActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func updateRepoActionVariableFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
value, err := params.GetString(req.GetArguments(), "value") value, err := params.GetString(args, "value")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -391,16 +391,16 @@ func updateRepoActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*
return to.TextResult(map[string]any{"message": "variable updated", "status": resp.StatusCode}) return to.TextResult(map[string]any{"message": "variable updated", "status": resp.StatusCode})
} }
func deleteRepoActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteRepoActionVariableFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -416,12 +416,12 @@ func deleteRepoActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*
return to.TextResult(map[string]any{"message": "variable deleted", "status": resp.StatusCode}) return to.TextResult(map[string]any{"message": "variable deleted", "status": resp.StatusCode})
} }
func listOrgActionVariablesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listOrgActionVariablesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
@@ -436,12 +436,12 @@ func listOrgActionVariablesFn(ctx context.Context, req mcp.CallToolRequest) (*mc
return to.TextResult(variables) return to.TextResult(variables)
} }
func getOrgActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getOrgActionVariableFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -457,20 +457,20 @@ func getOrgActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.
return to.TextResult(variable) return to.TextResult(variable)
} }
func createOrgActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createOrgActionVariableFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
value, err := params.GetString(req.GetArguments(), "value") value, err := params.GetString(args, "value")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
description, _ := req.GetArguments()["description"].(string) description, _ := args["description"].(string)
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
@@ -486,20 +486,20 @@ func createOrgActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(map[string]any{"message": "variable created", "status": resp.StatusCode}) return to.TextResult(map[string]any{"message": "variable created", "status": resp.StatusCode})
} }
func updateOrgActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func updateOrgActionVariableFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
value, err := params.GetString(req.GetArguments(), "value") value, err := params.GetString(args, "value")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
description, _ := req.GetArguments()["description"].(string) description, _ := args["description"].(string)
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
@@ -516,12 +516,12 @@ func updateOrgActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(map[string]any{"message": "variable updated", "status": resp.StatusCode}) return to.TextResult(map[string]any{"message": "variable updated", "status": resp.StatusCode})
} }
func deleteOrgActionVariableFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteOrgActionVariableFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
+107 -107
View File
@@ -14,9 +14,9 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/gitea" "gitea.com/gitea/gitea-mcp/pkg/gitea"
"gitea.com/gitea/gitea-mcp/pkg/params" "gitea.com/gitea/gitea-mcp/pkg/params"
"gitea.com/gitea/gitea-mcp/pkg/to" "gitea.com/gitea/gitea-mcp/pkg/to"
"gitea.com/gitea/gitea-mcp/pkg/tool"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
const ( const (
@@ -25,94 +25,94 @@ const (
) )
var ( var (
ActionsRunReadTool = mcp.NewTool( ActionsRunReadTool = tool.NewDefinition(
ActionsRunReadToolName, ActionsRunReadToolName,
mcp.WithDescription("Read Actions workflows, runs, jobs, logs, and artifacts."), "Read Actions workflows, runs, jobs, logs, and artifacts.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read Actions workflow, run, job, and artifact data")), annotation.ReadOnly("Read Actions workflow, run, job, and artifact data"),
mcp.WithString("method", mcp.Required(), mcp.Enum("list_workflows", "get_workflow", "list_runs", "get_run", "list_jobs", "list_run_jobs", "get_job", "get_job_log_preview", "download_job_log", "list_artifacts", "list_run_artifacts", "get_artifact", "download_artifact")), tool.String("method", tool.Required(), tool.Enum("list_workflows", "get_workflow", "list_runs", "get_run", "list_jobs", "list_run_jobs", "get_job", "get_job_log_preview", "download_job_log", "list_artifacts", "list_run_artifacts", "get_artifact", "download_artifact")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("workflow_id", mcp.Description("ID or filename (for 'get_workflow')")), tool.String("workflow_id", tool.Description("ID or filename (for 'get_workflow')")),
mcp.WithNumber("run_id", mcp.Description("for 'get_run'/'list_run_jobs'/'list_run_artifacts'")), tool.Number("run_id", tool.Description("for 'get_run'/'list_run_jobs'/'list_run_artifacts'")),
mcp.WithNumber("job_id", mcp.Description("for 'get_job'/log methods")), tool.Number("job_id", tool.Description("for 'get_job'/log methods")),
mcp.WithNumber("artifact_id", mcp.Description("for 'get_artifact'/'download_artifact'")), tool.Number("artifact_id", tool.Description("for 'get_artifact'/'download_artifact'")),
mcp.WithString("artifact_name", mcp.Description("name filter for 'list_artifacts'/'list_run_artifacts'")), tool.String("artifact_name", tool.Description("name filter for 'list_artifacts'/'list_run_artifacts'")),
mcp.WithString("status", mcp.Description("filter for 'list_runs'/'list_jobs'")), tool.String("status", tool.Description("filter for 'list_runs'/'list_jobs'")),
mcp.WithNumber("tail_lines", mcp.Description("log tail lines"), mcp.DefaultNumber(200), mcp.Min(1)), tool.Number("tail_lines", tool.Description("log tail lines"), tool.Default(200), tool.Minimum(1)),
mcp.WithNumber("max_bytes", mcp.Description("max log bytes"), mcp.DefaultNumber(65536), mcp.Min(1024)), tool.Number("max_bytes", tool.Description("max log bytes"), tool.Default(65536), tool.Minimum(1024)),
mcp.WithString("output_path", mcp.Description("for 'download_job_log'/'download_artifact'")), tool.String("output_path", tool.Description("for 'download_job_log'/'download_artifact'")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1), mcp.Min(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1), tool.Minimum(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30), mcp.Min(1)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30), tool.Minimum(1)),
) )
ActionsRunWriteTool = mcp.NewTool( ActionsRunWriteTool = tool.NewDefinition(
ActionsRunWriteToolName, ActionsRunWriteToolName,
mcp.WithDescription("Write Actions runs: dispatch, cancel, rerun."), "Write Actions runs: dispatch, cancel, rerun.",
mcp.WithToolAnnotation(annotation.Write("Trigger, cancel, or rerun Actions workflows")), annotation.Write("Trigger, cancel, or rerun Actions workflows"),
mcp.WithString("method", mcp.Required(), mcp.Enum("dispatch_workflow", "cancel_run", "rerun_run")), tool.String("method", tool.Required(), tool.Enum("dispatch_workflow", "cancel_run", "rerun_run")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("workflow_id", mcp.Description("ID or filename (for 'dispatch_workflow')")), tool.String("workflow_id", tool.Description("ID or filename (for 'dispatch_workflow')")),
mcp.WithString("ref", mcp.Description("branch or tag (for 'dispatch_workflow')")), tool.String("ref", tool.Description("branch or tag (for 'dispatch_workflow')")),
mcp.WithObject("inputs", mcp.Description("for 'dispatch_workflow'")), tool.Object("inputs", tool.Description("for 'dispatch_workflow'")),
mcp.WithNumber("run_id", mcp.Description("for 'cancel_run'/'rerun_run'")), tool.Number("run_id", tool.Description("for 'cancel_run'/'rerun_run'")),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{Tool: ActionsRunReadTool, Handler: runReadFn}) Tool.RegisterRead(tool.ServerTool{Tool: ActionsRunReadTool, Handler: runReadFn})
Tool.RegisterWrite(server.ServerTool{Tool: ActionsRunWriteTool, Handler: runWriteFn}) Tool.RegisterWrite(tool.ServerTool{Tool: ActionsRunWriteTool, Handler: runWriteFn})
} }
func runReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func runReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "list_workflows": case "list_workflows":
return listRepoActionWorkflowsFn(ctx, req) return listRepoActionWorkflowsFn(ctx, args)
case "get_workflow": case "get_workflow":
return getRepoActionWorkflowFn(ctx, req) return getRepoActionWorkflowFn(ctx, args)
case "list_runs": case "list_runs":
return listRepoActionRunsFn(ctx, req) return listRepoActionRunsFn(ctx, args)
case "get_run": case "get_run":
return getRepoActionRunFn(ctx, req) return getRepoActionRunFn(ctx, args)
case "list_jobs": case "list_jobs":
return listRepoActionJobsFn(ctx, req) return listRepoActionJobsFn(ctx, args)
case "list_run_jobs": case "list_run_jobs":
return listRepoActionRunJobsFn(ctx, req) return listRepoActionRunJobsFn(ctx, args)
case "get_job": case "get_job":
return getRepoActionJobFn(ctx, req) return getRepoActionJobFn(ctx, args)
case "get_job_log_preview": case "get_job_log_preview":
return getRepoActionJobLogPreviewFn(ctx, req) return getRepoActionJobLogPreviewFn(ctx, args)
case "download_job_log": case "download_job_log":
return downloadRepoActionJobLogFn(ctx, req) return downloadRepoActionJobLogFn(ctx, args)
case "list_artifacts": case "list_artifacts":
return listRepoActionArtifactsFn(ctx, req) return listRepoActionArtifactsFn(ctx, args)
case "list_run_artifacts": case "list_run_artifacts":
return listRepoActionRunArtifactsFn(ctx, req) return listRepoActionRunArtifactsFn(ctx, args)
case "get_artifact": case "get_artifact":
return getRepoActionArtifactFn(ctx, req) return getRepoActionArtifactFn(ctx, args)
case "download_artifact": case "download_artifact":
return downloadRepoActionArtifactFn(ctx, req) return downloadRepoActionArtifactFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func runWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func runWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "dispatch_workflow": case "dispatch_workflow":
return dispatchRepoActionWorkflowFn(ctx, req) return dispatchRepoActionWorkflowFn(ctx, args)
case "cancel_run": case "cancel_run":
return cancelRepoActionRunFn(ctx, req) return cancelRepoActionRunFn(ctx, args)
case "rerun_run": case "rerun_run":
return rerunRepoActionRunFn(ctx, req) return rerunRepoActionRunFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
@@ -135,16 +135,16 @@ func doJSONWithFallback(ctx context.Context, method string, paths []string, quer
return lastErr return lastErr
} }
func listRepoActionWorkflowsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoActionWorkflowsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
query := url.Values{} query := url.Values{}
query.Set("page", strconv.Itoa(page)) query.Set("page", strconv.Itoa(page))
query.Set("limit", strconv.Itoa(pageSize)) query.Set("limit", strconv.Itoa(pageSize))
@@ -162,16 +162,16 @@ func listRepoActionWorkflowsFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(slimActionWorkflows(result)) return to.TextResult(slimActionWorkflows(result))
} }
func getRepoActionWorkflowFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getRepoActionWorkflowFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
workflowID, err := params.GetString(req.GetArguments(), "workflow_id") workflowID, err := params.GetString(args, "workflow_id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -189,26 +189,26 @@ func getRepoActionWorkflowFn(ctx context.Context, req mcp.CallToolRequest) (*mcp
return to.TextResult(slimActionWorkflow(result)) return to.TextResult(slimActionWorkflow(result))
} }
func dispatchRepoActionWorkflowFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func dispatchRepoActionWorkflowFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
workflowID, err := params.GetString(req.GetArguments(), "workflow_id") workflowID, err := params.GetString(args, "workflow_id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
ref, err := params.GetString(req.GetArguments(), "ref") ref, err := params.GetString(args, "ref")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
var inputs map[string]any var inputs map[string]any
if raw, exists := req.GetArguments()["inputs"]; exists { if raw, exists := args["inputs"]; exists {
if m, ok := raw.(map[string]any); ok { if m, ok := raw.(map[string]any); ok {
inputs = m inputs = m
} }
@@ -238,17 +238,17 @@ func dispatchRepoActionWorkflowFn(ctx context.Context, req mcp.CallToolRequest)
return to.TextResult(map[string]any{"message": "workflow dispatched"}) return to.TextResult(map[string]any{"message": "workflow dispatched"})
} }
func listRepoActionRunsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoActionRunsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
statusFilter, _ := req.GetArguments()["status"].(string) statusFilter, _ := args["status"].(string)
query := url.Values{} query := url.Values{}
query.Set("page", strconv.Itoa(page)) query.Set("page", strconv.Itoa(page))
@@ -270,16 +270,16 @@ func listRepoActionRunsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Ca
return to.TextResult(slimActionRuns(result)) return to.TextResult(slimActionRuns(result))
} }
func getRepoActionRunFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getRepoActionRunFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
runID, err := params.GetIndex(req.GetArguments(), "run_id") runID, err := params.GetIndex(args, "run_id")
if err != nil || runID <= 0 { if err != nil || runID <= 0 {
return to.ErrorResult(errors.New("run_id is required")) return to.ErrorResult(errors.New("run_id is required"))
} }
@@ -297,16 +297,16 @@ func getRepoActionRunFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Call
return to.TextResult(slimActionRun(result)) return to.TextResult(slimActionRun(result))
} }
func cancelRepoActionRunFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func cancelRepoActionRunFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
runID, err := params.GetIndex(req.GetArguments(), "run_id") runID, err := params.GetIndex(args, "run_id")
if err != nil || runID <= 0 { if err != nil || runID <= 0 {
return to.ErrorResult(errors.New("run_id is required")) return to.ErrorResult(errors.New("run_id is required"))
} }
@@ -323,16 +323,16 @@ func cancelRepoActionRunFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.C
return to.TextResult(map[string]any{"message": "run cancellation requested"}) return to.TextResult(map[string]any{"message": "run cancellation requested"})
} }
func rerunRepoActionRunFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func rerunRepoActionRunFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
runID, err := params.GetIndex(req.GetArguments(), "run_id") runID, err := params.GetIndex(args, "run_id")
if err != nil || runID <= 0 { if err != nil || runID <= 0 {
return to.ErrorResult(errors.New("run_id is required")) return to.ErrorResult(errors.New("run_id is required"))
} }
@@ -354,17 +354,17 @@ func rerunRepoActionRunFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Ca
return to.TextResult(map[string]any{"message": "run rerun requested"}) return to.TextResult(map[string]any{"message": "run rerun requested"})
} }
func listRepoActionJobsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoActionJobsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
statusFilter, _ := req.GetArguments()["status"].(string) statusFilter, _ := args["status"].(string)
query := url.Values{} query := url.Values{}
query.Set("page", strconv.Itoa(page)) query.Set("page", strconv.Itoa(page))
@@ -386,20 +386,20 @@ func listRepoActionJobsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Ca
return to.TextResult(slimActionJobs(result)) return to.TextResult(slimActionJobs(result))
} }
func listRepoActionRunJobsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoActionRunJobsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
runID, err := params.GetIndex(req.GetArguments(), "run_id") runID, err := params.GetIndex(args, "run_id")
if err != nil || runID <= 0 { if err != nil || runID <= 0 {
return to.ErrorResult(errors.New("run_id is required")) return to.ErrorResult(errors.New("run_id is required"))
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
query := url.Values{} query := url.Values{}
query.Set("page", strconv.Itoa(page)) query.Set("page", strconv.Itoa(page))
@@ -418,16 +418,16 @@ func listRepoActionRunJobsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp
return to.TextResult(slimActionJobs(result)) return to.TextResult(slimActionJobs(result))
} }
func getRepoActionJobFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getRepoActionJobFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
jobID, err := params.GetIndex(req.GetArguments(), "job_id") jobID, err := params.GetIndex(args, "job_id")
if err != nil || jobID <= 0 { if err != nil || jobID <= 0 {
return to.ErrorResult(errors.New("job_id is required")) return to.ErrorResult(errors.New("job_id is required"))
} }
@@ -503,21 +503,21 @@ func limitBytes(data []byte, maxBytes int) ([]byte, bool) {
return data[len(data)-maxBytes:], true return data[len(data)-maxBytes:], true
} }
func getRepoActionJobLogPreviewFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getRepoActionJobLogPreviewFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
jobID, err := params.GetIndex(req.GetArguments(), "job_id") jobID, err := params.GetIndex(args, "job_id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
tailLines := int(params.GetOptionalInt(req.GetArguments(), "tail_lines", 200)) tailLines := int(params.GetOptionalInt(args, "tail_lines", 200))
maxBytes := int(params.GetOptionalInt(req.GetArguments(), "max_bytes", 65536)) maxBytes := int(params.GetOptionalInt(args, "max_bytes", 65536))
raw, usedPath, err := fetchJobLogBytes(ctx, owner, repo, jobID) raw, usedPath, err := fetchJobLogBytes(ctx, owner, repo, jobID)
if err != nil { if err != nil {
return to.ErrorResult(fmt.Errorf("get job log err: %v", err)) return to.ErrorResult(fmt.Errorf("get job log err: %v", err))
@@ -537,20 +537,20 @@ func getRepoActionJobLogPreviewFn(ctx context.Context, req mcp.CallToolRequest)
}) })
} }
func downloadRepoActionJobLogFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func downloadRepoActionJobLogFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
jobID, err := params.GetIndex(req.GetArguments(), "job_id") jobID, err := params.GetIndex(args, "job_id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
outputPath, _ := req.GetArguments()["output_path"].(string) outputPath, _ := args["output_path"].(string)
raw, usedPath, err := fetchJobLogBytes(ctx, owner, repo, jobID) raw, usedPath, err := fetchJobLogBytes(ctx, owner, repo, jobID)
if err != nil { if err != nil {
+41 -39
View File
@@ -3,7 +3,6 @@ package issue
import ( import (
"bytes" "bytes"
"context" "context"
"encoding/base64"
"errors" "errors"
"fmt" "fmt"
"io" "io"
@@ -17,51 +16,51 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/gitea" "gitea.com/gitea/gitea-mcp/pkg/gitea"
"gitea.com/gitea/gitea-mcp/pkg/params" "gitea.com/gitea/gitea-mcp/pkg/params"
"gitea.com/gitea/gitea-mcp/pkg/to" "gitea.com/gitea/gitea-mcp/pkg/to"
"gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
const AttachmentReadToolName = "attachment_read" const AttachmentReadToolName = "attachment_read"
var AttachmentReadTool = mcp.NewTool( var AttachmentReadTool = tool.NewDefinition(
AttachmentReadToolName, AttachmentReadToolName,
mcp.WithDescription("Read issue/comment attachments: list metadata, get metadata, or download content."), "Read issue/comment attachments: list metadata, get metadata, or download content.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read issue or comment attachments")), annotation.ReadOnly("Read issue or comment attachments"),
mcp.WithString("method", mcp.Required(), mcp.Enum("list", "get", "download")), tool.String("method", tool.Required(), tool.Enum("list", "get", "download")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("issue_number", mcp.Description("required for issue attachment list/get or issue-scoped metadata lookup")), tool.Number("issue_number", tool.Description("required for issue attachment list/get or issue-scoped metadata lookup")),
mcp.WithNumber("comment_id", mcp.Description("required for comment attachment list/get or comment-scoped metadata lookup")), tool.Number("comment_id", tool.Description("required for comment attachment list/get or comment-scoped metadata lookup")),
mcp.WithNumber("attachment_id", mcp.Description("required for get and for download when attachment_uuid is not provided")), tool.Number("attachment_id", tool.Description("required for get and for download when attachment_uuid is not provided")),
mcp.WithString("attachment_uuid", mcp.Description("attachment UUID for direct download path lookup")), tool.String("attachment_uuid", tool.Description("attachment UUID for direct download path lookup")),
mcp.WithString("output_path", mcp.Description("write the attachment to this exact path")), tool.String("output_path", tool.Description("write the attachment to this exact path")),
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{Tool: AttachmentReadTool, Handler: attachmentReadFn}) Tool.RegisterRead(tool.ServerTool{Tool: AttachmentReadTool, Handler: attachmentReadFn})
} }
func attachmentReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func attachmentReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "list": case "list":
return listAttachmentsFn(ctx, req) return listAttachmentsFn(ctx, args)
case "get": case "get":
return getAttachmentFn(ctx, req) return getAttachmentFn(ctx, args)
case "download": case "download":
return downloadAttachmentFn(ctx, req) return downloadAttachmentFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func listAttachmentsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listAttachmentsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, repo, issueNumber, commentID, err := attachmentScopeArgs(req) owner, repo, issueNumber, commentID, err := attachmentScopeArgs(args)
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -81,29 +80,29 @@ func listAttachmentsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
return to.TextResult(slimAttachments(attachments)) return to.TextResult(slimAttachments(attachments))
} }
func getAttachmentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getAttachmentFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
att, err := lookupAttachment(ctx, req) att, err := lookupAttachment(ctx, args)
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
return to.TextResult(slimAttachment(att)) return to.TextResult(slimAttachment(att))
} }
func downloadAttachmentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func downloadAttachmentFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
explicitOutputPath := params.GetOptionalString(req.GetArguments(), "output_path", "") explicitOutputPath := params.GetOptionalString(args, "output_path", "")
attachmentUUID := strings.TrimSpace(params.GetOptionalString(req.GetArguments(), "attachment_uuid", "")) attachmentUUID := strings.TrimSpace(params.GetOptionalString(args, "attachment_uuid", ""))
var att *gitea_sdk.Attachment var att *gitea_sdk.Attachment
if attachmentUUID == "" { if attachmentUUID == "" {
att, err = lookupAttachment(ctx, req) att, err = lookupAttachment(ctx, args)
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -132,7 +131,10 @@ func downloadAttachmentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Ca
} }
if len(limited) <= flag.MaxInlineAttachmentBytes { if len(limited) <= flag.MaxInlineAttachmentBytes {
text := fmt.Sprintf("attachment %s (%s, %d bytes, %s)", name, attachmentUUID, len(limited), mimeType) text := fmt.Sprintf("attachment %s (%s, %d bytes, %s)", name, attachmentUUID, len(limited), mimeType)
return mcp.NewToolResultImage(text, base64.StdEncoding.EncodeToString(limited), mimeType), nil return &mcp.CallToolResult{Content: []mcp.Content{
&mcp.TextContent{Text: text},
&mcp.ImageContent{Data: limited, MIMEType: mimeType},
}}, nil
} }
outputPath := defaultAttachmentPath(owner, repo, name, attachmentUUID) outputPath := defaultAttachmentPath(owner, repo, name, attachmentUUID)
if err := os.MkdirAll(filepath.Dir(outputPath), 0o700); err != nil { if err := os.MkdirAll(filepath.Dir(outputPath), 0o700); err != nil {
@@ -185,29 +187,29 @@ func attachmentFileResult(att *gitea_sdk.Attachment, outputPath string, written
return to.TextResult(res) return to.TextResult(res)
} }
func attachmentScopeArgs(req mcp.CallToolRequest) (owner, repo string, issueNumber, commentID int64, err error) { func attachmentScopeArgs(args map[string]any) (owner, repo string, issueNumber, commentID int64, err error) {
owner, err = params.GetString(req.GetArguments(), "owner") owner, err = params.GetString(args, "owner")
if err != nil { if err != nil {
return "", "", 0, 0, err return "", "", 0, 0, err
} }
repo, err = params.GetString(req.GetArguments(), "repo") repo, err = params.GetString(args, "repo")
if err != nil { if err != nil {
return "", "", 0, 0, err return "", "", 0, 0, err
} }
issueNumber = params.GetOptionalInt(req.GetArguments(), "issue_number", 0) issueNumber = params.GetOptionalInt(args, "issue_number", 0)
commentID = params.GetOptionalInt(req.GetArguments(), "comment_id", 0) commentID = params.GetOptionalInt(args, "comment_id", 0)
if (issueNumber > 0) == (commentID > 0) { if (issueNumber > 0) == (commentID > 0) {
return "", "", 0, 0, errors.New("exactly one of issue_number or comment_id is required") return "", "", 0, 0, errors.New("exactly one of issue_number or comment_id is required")
} }
return owner, repo, issueNumber, commentID, nil return owner, repo, issueNumber, commentID, nil
} }
func lookupAttachment(ctx context.Context, req mcp.CallToolRequest) (*gitea_sdk.Attachment, error) { func lookupAttachment(ctx context.Context, args map[string]any) (*gitea_sdk.Attachment, error) {
owner, repo, issueNumber, commentID, err := attachmentScopeArgs(req) owner, repo, issueNumber, commentID, err := attachmentScopeArgs(args)
if err != nil { if err != nil {
return nil, err return nil, err
} }
attachmentID := params.GetOptionalInt(req.GetArguments(), "attachment_id", 0) attachmentID := params.GetOptionalInt(args, "attachment_id", 0)
if attachmentID <= 0 { if attachmentID <= 0 {
return nil, errors.New("attachment_id is required") return nil, errors.New("attachment_id is required")
} }
+66 -7
View File
@@ -1,7 +1,9 @@
package issue package issue
import ( import (
"bytes"
"context" "context"
"encoding/base64"
"encoding/json" "encoding/json"
"fmt" "fmt"
"net/http" "net/http"
@@ -14,7 +16,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"gitea.com/gitea/gitea-mcp/pkg/gitea" "gitea.com/gitea/gitea-mcp/pkg/gitea"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
func TestAttachmentFilename(t *testing.T) { func TestAttachmentFilename(t *testing.T) {
@@ -79,13 +81,13 @@ func TestAttachmentReadListIssueAttachments(t *testing.T) {
flag.Host, flag.Token, flag.Version = server.URL, "", "test" flag.Host, flag.Token, flag.Version = server.URL, "", "test"
defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }() defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }()
res, err := attachmentReadFn(context.Background(), mcp.CallToolRequest{Params: mcp.CallToolParams{Arguments: map[string]any{ res, err := attachmentReadFn(context.Background(), map[string]any{
"method": "list", "owner": owner, "repo": repo, "issue_number": float64(42), "method": "list", "owner": owner, "repo": repo, "issue_number": float64(42),
}}}) })
if err != nil { if err != nil {
t.Fatalf("attachmentReadFn() error = %v", err) t.Fatalf("attachmentReadFn() error = %v", err)
} }
body := res.Content[0].(mcp.TextContent).Text body := res.Content[0].(*mcp.TextContent).Text
if !strings.Contains(body, `"mime_type":"image/png"`) || !strings.Contains(body, `"uuid":"uuid-1"`) { if !strings.Contains(body, `"mime_type":"image/png"`) || !strings.Contains(body, `"uuid":"uuid-1"`) {
t.Fatalf("unexpected body: %s", body) t.Fatalf("unexpected body: %s", body)
} }
@@ -158,13 +160,13 @@ func TestAttachmentReadDownloadSavesLargeAttachmentToDefaultFile(t *testing.T) {
flag.Host, flag.Token, flag.Version, flag.MaxInlineAttachmentBytes = origHost, origToken, origVersion, origInline flag.Host, flag.Token, flag.Version, flag.MaxInlineAttachmentBytes = origHost, origToken, origVersion, origInline
}() }()
res, err := attachmentReadFn(context.Background(), mcp.CallToolRequest{Params: mcp.CallToolParams{Arguments: map[string]any{ res, err := attachmentReadFn(context.Background(), map[string]any{
"method": "download", "owner": owner, "repo": repo, "issue_number": float64(42), "attachment_id": float64(1), "method": "download", "owner": owner, "repo": repo, "issue_number": float64(42), "attachment_id": float64(1),
}}}) })
if err != nil { if err != nil {
t.Fatalf("attachmentReadFn() error = %v", err) t.Fatalf("attachmentReadFn() error = %v", err)
} }
body := res.Content[0].(mcp.TextContent).Text body := res.Content[0].(*mcp.TextContent).Text
wantPath := filepath.Join(home, ".gitea-mcp", "attachments", owner, repo, "large-uuid-1.bin") wantPath := filepath.Join(home, ".gitea-mcp", "attachments", owner, repo, "large-uuid-1.bin")
if !strings.Contains(body, wantPath) { if !strings.Contains(body, wantPath) {
t.Fatalf("result missing path %q: %s", wantPath, body) t.Fatalf("result missing path %q: %s", wantPath, body)
@@ -180,3 +182,60 @@ func TestAttachmentReadDownloadSavesLargeAttachmentToDefaultFile(t *testing.T) {
t.Fatalf("result missing bytes: %s", body) t.Fatalf("result missing bytes: %s", body)
} }
} }
func TestAttachmentReadDownloadReturnsRawImageContent(t *testing.T) {
const uuid = "uuid-1"
payload := []byte{0, 1, 2, 250}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/attachments/"+uuid {
http.NotFound(w, r)
return
}
w.Header().Set("Content-Type", "image/png")
_, _ = w.Write(payload)
}))
defer server.Close()
originalHost := flag.Host
originalLimit := flag.MaxInlineAttachmentBytes
flag.Host = server.URL
flag.MaxInlineAttachmentBytes = len(payload)
defer func() {
flag.Host = originalHost
flag.MaxInlineAttachmentBytes = originalLimit
}()
result, err := attachmentReadFn(context.Background(), map[string]any{
"method": "download",
"owner": "octo",
"repo": "demo",
"attachment_uuid": uuid,
})
if err != nil {
t.Fatalf("attachmentReadFn() error = %v", err)
}
if len(result.Content) != 2 {
t.Fatalf("content count = %d, want 2", len(result.Content))
}
if _, ok := result.Content[0].(*mcp.TextContent); !ok {
t.Fatalf("first content type = %T, want *mcp.TextContent", result.Content[0])
}
image, ok := result.Content[1].(*mcp.ImageContent)
if !ok {
t.Fatalf("second content type = %T, want *mcp.ImageContent", result.Content[1])
}
if image.MIMEType != "image/png" {
t.Errorf("image MIME type = %q, want image/png", image.MIMEType)
}
if !bytes.Equal(image.Data, payload) {
t.Errorf("image data = %v, want raw payload %v", image.Data, payload)
}
wire, err := json.Marshal(image)
if err != nil {
t.Fatalf("json.Marshal() error = %v", err)
}
wantBase64 := base64.StdEncoding.EncodeToString(payload)
if !strings.Contains(string(wire), `"data":"`+wantBase64+`"`) {
t.Errorf("wire image = %s, want base64 data %q", wire, wantBase64)
}
}
+120 -124
View File
@@ -13,8 +13,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
// issueWithAssets / commentWithAssets wrap the SDK types to capture the // issueWithAssets / commentWithAssets wrap the SDK types to capture the
@@ -38,125 +37,123 @@ const (
) )
var ( var (
ListRepoIssuesTool = mcp.NewTool( ListRepoIssuesTool = tool.NewDefinition(
ListRepoIssuesToolName, ListRepoIssuesToolName,
mcp.WithDescription("List issues in a repository (or pull requests, via the 'type' filter), filterable by state, labels, milestones, and update time range."), "List issues in a repository (or pull requests, via the 'type' filter), filterable by state, labels, milestones, and update time range.",
mcp.WithToolAnnotation(annotation.ReadOnly("List repository issues")), annotation.ReadOnly("List repository issues"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("state", mcp.DefaultString("all")), tool.String("state", tool.Default("all")),
mcp.WithString("type", mcp.Description("issues or pulls"), mcp.Enum("issues", "pulls")), tool.String("type", tool.Description("issues or pulls"), tool.Enum("issues", "pulls")),
mcp.WithArray("labels", mcp.Description("label name filter"), mcp.Items(map[string]any{"type": "string"})), tool.Array("labels", tool.Description("label name filter"), tool.Items(map[string]any{"type": "string"})),
mcp.WithArray("milestones", mcp.Description("milestone name or ID filter"), mcp.Items(map[string]any{"type": "string"})), tool.Array("milestones", tool.Description("milestone name or ID filter"), tool.Items(map[string]any{"type": "string"})),
mcp.WithString("since", mcp.Description("updated after ISO 8601")), tool.String("since", tool.Description("updated after ISO 8601")),
mcp.WithString("before", mcp.Description("updated before ISO 8601")), tool.String("before", tool.Description("updated before ISO 8601")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
IssueReadTool = mcp.NewTool( IssueReadTool = tool.NewDefinition(
IssueReadToolName, IssueReadToolName,
mcp.WithDescription("Read issue: details, comments, or labels."), "Read issue: details, comments, or labels.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read issue details")), annotation.ReadOnly("Read issue details"),
mcp.WithString("method", mcp.Required(), mcp.Enum("get", "get_comments", "get_labels")), tool.String("method", tool.Required(), tool.Enum("get", "get_comments", "get_labels")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("issue_number", mcp.Required()), tool.Number("issue_number", tool.Required()),
) )
IssueWriteTool = mcp.NewTool( IssueWriteTool = tool.NewDefinition(
IssueWriteToolName, IssueWriteToolName,
mcp.WithDescription("Write issues: create, update, manage comments and labels."), "Write issues: create, update, manage comments and labels.",
mcp.WithToolAnnotation(annotation.Write("Create or update issues, comments, and labels")), annotation.Write("Create or update issues, comments, and labels"),
mcp.WithString("method", mcp.Required(), mcp.Enum("create", "update", "add_comment", "edit_comment", "add_labels", "remove_label", "replace_labels", "clear_labels")), tool.String("method", tool.Required(), tool.Enum("create", "update", "add_comment", "edit_comment", "add_labels", "remove_label", "replace_labels", "clear_labels")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("issue_number", mcp.Description("required except for 'create'")), tool.Number("issue_number", tool.Description("required except for 'create'")),
mcp.WithString("title", mcp.Description("required for 'create'")), tool.String("title", tool.Description("required for 'create'")),
mcp.WithString("body", mcp.Description("required for 'create'/'add_comment'/'edit_comment'")), tool.String("body", tool.Description("required for 'create'/'add_comment'/'edit_comment'")),
mcp.WithArray("assignees", mcp.Items(map[string]any{"type": "string"})), tool.Array("assignees", tool.Items(map[string]any{"type": "string"})),
mcp.WithNumber("milestone"), tool.Number("milestone"),
mcp.WithString("state", mcp.Enum("open", "closed", "all")), tool.String("state", tool.Enum("open", "closed", "all")),
mcp.WithNumber("commentID", mcp.Description("for 'edit_comment'")), tool.Number("commentID", tool.Description("for 'edit_comment'")),
mcp.WithArray("labels", mcp.Description("label IDs"), mcp.Items(map[string]any{"type": "number"})), tool.Array("labels", tool.Description("label IDs"), tool.Items(map[string]any{"type": "number"})),
mcp.WithNumber("label_id", mcp.Description("for 'remove_label'")), tool.Number("label_id", tool.Description("for 'remove_label'")),
mcp.WithString("ref", mcp.Description("branch to associate")), tool.String("ref", tool.Description("branch to associate")),
mcp.WithString("deadline", mcp.Description("ISO 8601")), tool.String("deadline", tool.Description("ISO 8601")),
mcp.WithBoolean("remove_deadline"), tool.Boolean("remove_deadline"),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: ListRepoIssuesTool, Tool: ListRepoIssuesTool,
Handler: listRepoIssuesFn, Handler: listRepoIssuesFn,
}) })
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: IssueReadTool, Tool: IssueReadTool,
Handler: issueReadFn, Handler: issueReadFn,
}) })
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: IssueWriteTool, Tool: IssueWriteTool,
Handler: issueWriteFn, Handler: issueWriteFn,
}) })
} }
func issueReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func issueReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
method, err := params.GetString(args, "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "get": case "get":
return getIssueByIndexFn(ctx, req) return getIssueByIndexFn(ctx, args)
case "get_comments": case "get_comments":
return getIssueCommentsByIndexFn(ctx, req) return getIssueCommentsByIndexFn(ctx, args)
case "get_labels": case "get_labels":
return getIssueLabelsFn(ctx, req) return getIssueLabelsFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func issueWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func issueWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
method, err := params.GetString(args, "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "create": case "create":
return createIssueFn(ctx, req) return createIssueFn(ctx, args)
case "update": case "update":
return editIssueFn(ctx, req) return editIssueFn(ctx, args)
case "add_comment": case "add_comment":
return createIssueCommentFn(ctx, req) return createIssueCommentFn(ctx, args)
case "edit_comment": case "edit_comment":
return editIssueCommentFn(ctx, req) return editIssueCommentFn(ctx, args)
case "add_labels": case "add_labels":
return addIssueLabelsFn(ctx, req) return addIssueLabelsFn(ctx, args)
case "remove_label": case "remove_label":
return removeIssueLabelFn(ctx, req) return removeIssueLabelFn(ctx, args)
case "replace_labels": case "replace_labels":
return replaceIssueLabelsFn(ctx, req) return replaceIssueLabelsFn(ctx, args)
case "clear_labels": case "clear_labels":
return clearIssueLabelsFn(ctx, req) return clearIssueLabelsFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func getIssueByIndexFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getIssueByIndexFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -170,22 +167,22 @@ func getIssueByIndexFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
return to.TextResult(m) return to.TextResult(m)
} }
func listRepoIssuesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoIssuesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
state, ok := req.GetArguments()["state"].(string) state, ok := args["state"].(string)
if !ok { if !ok {
state = "all" state = "all"
} }
labels := params.GetStringSlice(req.GetArguments(), "labels") labels := params.GetStringSlice(args, "labels")
milestones := params.GetStringSlice(req.GetArguments(), "milestones") milestones := params.GetStringSlice(args, "milestones")
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
opt := gitea_sdk.ListIssueOption{ opt := gitea_sdk.ListIssueOption{
State: gitea_sdk.StateType(state), State: gitea_sdk.StateType(state),
Labels: labels, Labels: labels,
@@ -195,16 +192,16 @@ func listRepoIssuesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
PageSize: pageSize, PageSize: pageSize,
}, },
} }
switch req.GetArguments()["type"] { switch args["type"] {
case "issues": case "issues":
opt.Type = gitea_sdk.IssueTypeIssue opt.Type = gitea_sdk.IssueTypeIssue
case "pulls": case "pulls":
opt.Type = gitea_sdk.IssueTypePull opt.Type = gitea_sdk.IssueTypePull
} }
if t := params.GetOptionalTime(req.GetArguments(), "since"); t != nil { if t := params.GetOptionalTime(args, "since"); t != nil {
opt.Since = *t opt.Since = *t
} }
if t := params.GetOptionalTime(req.GetArguments(), "before"); t != nil { if t := params.GetOptionalTime(args, "before"); t != nil {
opt.Before = *t opt.Before = *t
} }
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
@@ -218,20 +215,20 @@ func listRepoIssuesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(slimIssues(issues)) return to.TextResult(slimIssues(issues))
} }
func createIssueFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createIssueFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
title, err := params.GetString(req.GetArguments(), "title") title, err := params.GetString(args, "title")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
body, err := params.GetString(req.GetArguments(), "body") body, err := params.GetString(args, "body")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -243,19 +240,19 @@ func createIssueFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolR
Title: title, Title: title,
Body: body, Body: body,
} }
opt.Assignees = params.GetStringSlice(req.GetArguments(), "assignees") opt.Assignees = params.GetStringSlice(args, "assignees")
if val, exists := req.GetArguments()["milestone"]; exists { if val, exists := args["milestone"]; exists {
if milestone, ok := params.ToInt64(val); ok { if milestone, ok := params.ToInt64(val); ok {
opt.Milestone = milestone opt.Milestone = milestone
} }
} }
if labelIDs, err := params.GetInt64Slice(req.GetArguments(), "labels"); err == nil { if labelIDs, err := params.GetInt64Slice(args, "labels"); err == nil {
opt.Labels = labelIDs opt.Labels = labelIDs
} }
if ref, ok := req.GetArguments()["ref"].(string); ok { if ref, ok := args["ref"].(string); ok {
opt.Ref = ref opt.Ref = ref
} }
opt.Deadline = params.GetOptionalTime(req.GetArguments(), "deadline") opt.Deadline = params.GetOptionalTime(args, "deadline")
issue, _, err := client.Issues.CreateIssue(ctx, owner, repo, opt) issue, _, err := client.Issues.CreateIssue(ctx, owner, repo, opt)
if err != nil { if err != nil {
return to.ErrorResult(fmt.Errorf("create %v/%v/issue err: %v", owner, repo, err)) return to.ErrorResult(fmt.Errorf("create %v/%v/issue err: %v", owner, repo, err))
@@ -264,20 +261,20 @@ func createIssueFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolR
return to.TextResult(slimIssue(issue)) return to.TextResult(slimIssue(issue))
} }
func createIssueCommentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createIssueCommentFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
body, err := params.GetString(req.GetArguments(), "body") body, err := params.GetString(args, "body")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -296,21 +293,20 @@ func createIssueCommentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Ca
return to.TextResult(slimComment(issueComment)) return to.TextResult(slimComment(issueComment))
} }
func editIssueFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func editIssueFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
args := req.GetArguments()
opt := gitea_sdk.EditIssueOption{ opt := gitea_sdk.EditIssueOption{
Body: params.GetPresentStringPtr(args, "body"), Body: params.GetPresentStringPtr(args, "body"),
Ref: params.GetPresentStringPtr(args, "ref"), Ref: params.GetPresentStringPtr(args, "ref"),
@@ -343,20 +339,20 @@ func editIssueFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolRes
return to.TextResult(slimIssue(issue)) return to.TextResult(slimIssue(issue))
} }
func editIssueCommentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func editIssueCommentFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
commentID, err := params.GetIndex(req.GetArguments(), "commentID") commentID, err := params.GetIndex(args, "commentID")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
body, err := params.GetString(req.GetArguments(), "body") body, err := params.GetString(args, "body")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -375,16 +371,16 @@ func editIssueCommentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Call
return to.TextResult(slimComment(issueComment)) return to.TextResult(slimComment(issueComment))
} }
func getIssueCommentsByIndexFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getIssueCommentsByIndexFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -402,16 +398,16 @@ func getIssueCommentsByIndexFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(out) return to.TextResult(out)
} }
func getIssueLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getIssueLabelsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -427,20 +423,20 @@ func getIssueLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(slim.Labels(labels)) return to.TextResult(slim.Labels(labels))
} }
func addIssueLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func addIssueLabelsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
labels, err := params.GetInt64Slice(req.GetArguments(), "labels") labels, err := params.GetInt64Slice(args, "labels")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -456,20 +452,20 @@ func addIssueLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(slim.Labels(issueLabels)) return to.TextResult(slim.Labels(issueLabels))
} }
func replaceIssueLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func replaceIssueLabelsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
labels, err := params.GetInt64Slice(req.GetArguments(), "labels") labels, err := params.GetInt64Slice(args, "labels")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -485,16 +481,16 @@ func replaceIssueLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Ca
return to.TextResult(slim.Labels(issueLabels)) return to.TextResult(slim.Labels(issueLabels))
} }
func clearIssueLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func clearIssueLabelsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -510,20 +506,20 @@ func clearIssueLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Call
return to.TextResult("Labels cleared successfully") return to.TextResult("Labels cleared successfully")
} }
func removeIssueLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func removeIssueLabelFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
labelID, err := params.GetIndex(req.GetArguments(), "label_id") labelID, err := params.GetIndex(args, "label_id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
+29 -37
View File
@@ -12,7 +12,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
func Test_listRepoIssuesFn_filters(t *testing.T) { func Test_listRepoIssuesFn_filters(t *testing.T) {
@@ -60,20 +60,16 @@ func Test_listRepoIssuesFn_filters(t *testing.T) {
flag.Version = origVersion flag.Version = origVersion
}() }()
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "owner": owner,
Arguments: map[string]any{ "repo": repo,
"owner": owner, "type": "issues",
"repo": repo, "labels": []any{"bug", "enhancement"},
"type": "issues", "milestones": []any{"v1.0", "2"},
"labels": []any{"bug", "enhancement"}, "since": "2026-01-01T00:00:00Z",
"milestones": []any{"v1.0", "2"},
"since": "2026-01-01T00:00:00Z",
},
},
} }
_, err := listRepoIssuesFn(context.Background(), req) _, err := listRepoIssuesFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("listRepoIssuesFn() error = %v", err) t.Fatalf("listRepoIssuesFn() error = %v", err)
} }
@@ -126,17 +122,17 @@ func Test_listRepoIssuesFn_includesMilestone(t *testing.T) {
flag.Host, flag.Token, flag.Version = server.URL, "", "test" flag.Host, flag.Token, flag.Version = server.URL, "", "test"
defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }() defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }()
req := mcp.CallToolRequest{Params: mcp.CallToolParams{Arguments: map[string]any{ args := map[string]any{
"owner": owner, "repo": repo, "owner": owner, "repo": repo,
}}} }
res, err := listRepoIssuesFn(context.Background(), req) res, err := listRepoIssuesFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("listRepoIssuesFn() error = %v", err) t.Fatalf("listRepoIssuesFn() error = %v", err)
} }
if res.IsError { if res.IsError {
t.Fatalf("unexpected error result: %v", res.Content) t.Fatalf("unexpected error result: %v", res.Content)
} }
body := res.Content[0].(mcp.TextContent).Text body := res.Content[0].(*mcp.TextContent).Text
if !strings.Contains(body, `"milestone"`) || !strings.Contains(body, `"v1.0"`) { if !strings.Contains(body, `"milestone"`) || !strings.Contains(body, `"v1.0"`) {
t.Fatalf("expected milestone in list output, got: %s", body) t.Fatalf("expected milestone in list output, got: %s", body)
} }
@@ -189,20 +185,16 @@ func Test_createIssueFn_labels(t *testing.T) {
flag.Version = origVersion flag.Version = origVersion
}() }()
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "owner": owner,
Arguments: map[string]any{ "repo": repo,
"owner": owner, "title": "test issue",
"repo": repo, "body": "body",
"title": "test issue", "labels": []any{float64(10), float64(20)},
"body": "body", "deadline": "2026-06-01T00:00:00Z",
"labels": []any{float64(10), float64(20)},
"deadline": "2026-06-01T00:00:00Z",
},
},
} }
_, err := createIssueFn(context.Background(), req) _, err := createIssueFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("createIssueFn() error = %v", err) t.Fatalf("createIssueFn() error = %v", err)
} }
@@ -255,17 +247,17 @@ func Test_getIssueByIndexFn_includesAttachments(t *testing.T) {
flag.Host, flag.Token, flag.Version = server.URL, "", "test" flag.Host, flag.Token, flag.Version = server.URL, "", "test"
defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }() defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }()
req := mcp.CallToolRequest{Params: mcp.CallToolParams{Arguments: map[string]any{ args := map[string]any{
"owner": owner, "repo": repo, "issue_number": float64(42), "owner": owner, "repo": repo, "issue_number": float64(42),
}}} }
res, err := getIssueByIndexFn(context.Background(), req) res, err := getIssueByIndexFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("getIssueByIndexFn() error = %v", err) t.Fatalf("getIssueByIndexFn() error = %v", err)
} }
if res.IsError { if res.IsError {
t.Fatalf("unexpected error result: %v", res.Content) t.Fatalf("unexpected error result: %v", res.Content)
} }
body := res.Content[0].(mcp.TextContent).Text body := res.Content[0].(*mcp.TextContent).Text
if !strings.Contains(body, `[shot.png](https://example/shot.png)`) { if !strings.Contains(body, `[shot.png](https://example/shot.png)`) {
t.Fatalf("expected attachment markdown inlined in body, got: %s", body) t.Fatalf("expected attachment markdown inlined in body, got: %s", body)
} }
@@ -304,17 +296,17 @@ func Test_getIssueCommentsByIndexFn_includesAttachments(t *testing.T) {
flag.Host, flag.Token, flag.Version = server.URL, "", "test" flag.Host, flag.Token, flag.Version = server.URL, "", "test"
defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }() defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }()
req := mcp.CallToolRequest{Params: mcp.CallToolParams{Arguments: map[string]any{ args := map[string]any{
"owner": owner, "repo": repo, "issue_number": float64(7), "owner": owner, "repo": repo, "issue_number": float64(7),
}}} }
res, err := getIssueCommentsByIndexFn(context.Background(), req) res, err := getIssueCommentsByIndexFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("getIssueCommentsByIndexFn() error = %v", err) t.Fatalf("getIssueCommentsByIndexFn() error = %v", err)
} }
if res.IsError { if res.IsError {
t.Fatalf("unexpected error result: %v", res.Content) t.Fatalf("unexpected error result: %v", res.Content)
} }
body := res.Content[0].(mcp.TextContent).Text body := res.Content[0].(*mcp.TextContent).Text
if !strings.Contains(body, `[log.txt](https://example/log.txt)`) { if !strings.Contains(body, `[log.txt](https://example/log.txt)`) {
t.Fatalf("expected attachment markdown inlined in body, got: %s", body) t.Fatalf("expected attachment markdown inlined in body, got: %s", body)
} }
+75 -80
View File
@@ -12,8 +12,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("label") var Tool = tool.New("label")
@@ -24,99 +23,97 @@ const (
) )
var ( var (
LabelReadTool = mcp.NewTool( LabelReadTool = tool.NewDefinition(
LabelReadToolName, LabelReadToolName,
mcp.WithDescription("Read repo or org labels."), "Read repo or org labels.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read labels")), annotation.ReadOnly("Read labels"),
mcp.WithString("method", mcp.Required(), mcp.Enum("list_repo_labels", "get_repo_label", "list_org_labels")), tool.String("method", tool.Required(), tool.Enum("list_repo_labels", "get_repo_label", "list_org_labels")),
mcp.WithString("owner", mcp.Description("for repo methods")), tool.String("owner", tool.Description("for repo methods")),
mcp.WithString("repo", mcp.Description("for repo methods")), tool.String("repo", tool.Description("for repo methods")),
mcp.WithString("org", mcp.Description("for org methods")), tool.String("org", tool.Description("for org methods")),
mcp.WithNumber("id", mcp.Description("label ID (for 'get_repo_label')")), tool.Number("id", tool.Description("label ID (for 'get_repo_label')")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
LabelWriteTool = mcp.NewTool( LabelWriteTool = tool.NewDefinition(
LabelWriteToolName, LabelWriteToolName,
mcp.WithDescription("Write labels (repo or org): create, edit, delete."), "Write labels (repo or org): create, edit, delete.",
mcp.WithToolAnnotation(annotation.Destructive("Create, update, or delete labels")), annotation.Destructive("Create, update, or delete labels"),
mcp.WithString("method", mcp.Required(), mcp.Enum("create_repo_label", "edit_repo_label", "delete_repo_label", "create_org_label", "edit_org_label", "delete_org_label")), tool.String("method", tool.Required(), tool.Enum("create_repo_label", "edit_repo_label", "delete_repo_label", "create_org_label", "edit_org_label", "delete_org_label")),
mcp.WithString("owner", mcp.Description("for repo methods")), tool.String("owner", tool.Description("for repo methods")),
mcp.WithString("repo", mcp.Description("for repo methods")), tool.String("repo", tool.Description("for repo methods")),
mcp.WithString("org", mcp.Description("for org methods")), tool.String("org", tool.Description("for org methods")),
mcp.WithNumber("id", mcp.Description("for edit/delete")), tool.Number("id", tool.Description("for edit/delete")),
mcp.WithString("name", mcp.Description("required for create")), tool.String("name", tool.Description("required for create")),
mcp.WithString("color", mcp.Description("hex (#RRGGBB); required for create")), tool.String("color", tool.Description("hex (#RRGGBB); required for create")),
mcp.WithString("description"), tool.String("description"),
mcp.WithBoolean("exclusive", mcp.Description("exclusive (org only)")), tool.Boolean("exclusive", tool.Description("exclusive (org only)")),
mcp.WithBoolean("is_archived", mcp.Description("archived (repo only)")), tool.Boolean("is_archived", tool.Description("archived (repo only)")),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: LabelReadTool, Tool: LabelReadTool,
Handler: labelReadFn, Handler: labelReadFn,
}) })
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: LabelWriteTool, Tool: LabelWriteTool,
Handler: labelWriteFn, Handler: labelWriteFn,
}) })
} }
func labelReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func labelReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
method, err := params.GetString(args, "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "list_repo_labels": case "list_repo_labels":
return listRepoLabelsFn(ctx, req) return listRepoLabelsFn(ctx, args)
case "get_repo_label": case "get_repo_label":
return getRepoLabelFn(ctx, req) return getRepoLabelFn(ctx, args)
case "list_org_labels": case "list_org_labels":
return listOrgLabelsFn(ctx, req) return listOrgLabelsFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func labelWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func labelWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
method, err := params.GetString(args, "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "create_repo_label": case "create_repo_label":
return createRepoLabelFn(ctx, req) return createRepoLabelFn(ctx, args)
case "edit_repo_label": case "edit_repo_label":
return editRepoLabelFn(ctx, req) return editRepoLabelFn(ctx, args)
case "delete_repo_label": case "delete_repo_label":
return deleteRepoLabelFn(ctx, req) return deleteRepoLabelFn(ctx, args)
case "create_org_label": case "create_org_label":
return createOrgLabelFn(ctx, req) return createOrgLabelFn(ctx, args)
case "edit_org_label": case "edit_org_label":
return editOrgLabelFn(ctx, req) return editOrgLabelFn(ctx, args)
case "delete_org_label": case "delete_org_label":
return deleteOrgLabelFn(ctx, req) return deleteOrgLabelFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func listRepoLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoLabelsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
opt := gitea_sdk.ListLabelsOptions{ opt := gitea_sdk.ListLabelsOptions{
ListOptions: gitea_sdk.ListOptions{ ListOptions: gitea_sdk.ListOptions{
@@ -135,16 +132,16 @@ func listRepoLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(slim.Labels(labels)) return to.TextResult(slim.Labels(labels))
} }
func getRepoLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getRepoLabelFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -160,26 +157,26 @@ func getRepoLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTool
return to.TextResult(slim.Label(label)) return to.TextResult(slim.Label(label))
} }
func createRepoLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createRepoLabelFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
color, err := params.GetString(req.GetArguments(), "color") color, err := params.GetString(args, "color")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
description, _ := req.GetArguments()["description"].(string) // Optional description, _ := args["description"].(string) // Optional
isArchived, _ := req.GetArguments()["is_archived"].(bool) isArchived, _ := args["is_archived"].(bool)
opt := gitea_sdk.CreateLabelOption{ opt := gitea_sdk.CreateLabelOption{
Name: name, Name: name,
@@ -199,21 +196,20 @@ func createRepoLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
return to.TextResult(slim.Label(label)) return to.TextResult(slim.Label(label))
} }
func editRepoLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func editRepoLabelFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
args := req.GetArguments()
opt := gitea_sdk.EditLabelOption{ opt := gitea_sdk.EditLabelOption{
Name: params.GetOptionalStringPtr(args, "name"), Name: params.GetOptionalStringPtr(args, "name"),
Color: params.GetOptionalStringPtr(args, "color"), Color: params.GetOptionalStringPtr(args, "color"),
@@ -232,16 +228,16 @@ func editRepoLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToo
return to.TextResult(slim.Label(label)) return to.TextResult(slim.Label(label))
} }
func deleteRepoLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteRepoLabelFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -257,12 +253,12 @@ func deleteRepoLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
return to.TextResult("Label deleted successfully") return to.TextResult("Label deleted successfully")
} }
func listOrgLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listOrgLabelsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
opt := gitea_sdk.ListOrgLabelsOptions{ opt := gitea_sdk.ListOrgLabelsOptions{
ListOptions: gitea_sdk.ListOptions{ ListOptions: gitea_sdk.ListOptions{
@@ -281,21 +277,21 @@ func listOrgLabelsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToo
return to.TextResult(slim.Labels(labels)) return to.TextResult(slim.Labels(labels))
} }
func createOrgLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createOrgLabelFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
name, err := params.GetString(req.GetArguments(), "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
color, err := params.GetString(req.GetArguments(), "color") color, err := params.GetString(args, "color")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
description, _ := req.GetArguments()["description"].(string) description, _ := args["description"].(string)
exclusive, _ := req.GetArguments()["exclusive"].(bool) exclusive, _ := args["exclusive"].(bool)
opt := gitea_sdk.CreateOrgLabelOption{ opt := gitea_sdk.CreateOrgLabelOption{
Name: name, Name: name,
@@ -315,17 +311,16 @@ func createOrgLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(slim.Label(label)) return to.TextResult(slim.Label(label))
} }
func editOrgLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func editOrgLabelFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
args := req.GetArguments()
opt := gitea_sdk.EditOrgLabelOption{ opt := gitea_sdk.EditOrgLabelOption{
Name: params.GetOptionalStringPtr(args, "name"), Name: params.GetOptionalStringPtr(args, "name"),
Color: params.GetOptionalStringPtr(args, "color"), Color: params.GetOptionalStringPtr(args, "color"),
@@ -344,12 +339,12 @@ func editOrgLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTool
return to.TextResult(slim.Label(label)) return to.TextResult(slim.Label(label))
} }
func deleteOrgLabelFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteOrgLabelFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
+59 -61
View File
@@ -11,8 +11,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("milestone") var Tool = tool.New("milestone")
@@ -23,90 +22,90 @@ const (
) )
var ( var (
MilestoneReadTool = mcp.NewTool( MilestoneReadTool = tool.NewDefinition(
MilestoneReadToolName, MilestoneReadToolName,
mcp.WithDescription("Read milestones: get one or list."), "Read milestones: get one or list.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read milestones")), annotation.ReadOnly("Read milestones"),
mcp.WithString("method", mcp.Required(), mcp.Enum("get", "list")), tool.String("method", tool.Required(), tool.Enum("get", "list")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("id", mcp.Description("for 'get'")), tool.Number("id", tool.Description("for 'get'")),
mcp.WithString("state", mcp.DefaultString("all")), tool.String("state", tool.Default("all")),
mcp.WithString("name", mcp.Description("name filter (for 'list')")), tool.String("name", tool.Description("name filter (for 'list')")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
MilestoneWriteTool = mcp.NewTool( MilestoneWriteTool = tool.NewDefinition(
MilestoneWriteToolName, MilestoneWriteToolName,
mcp.WithDescription("Write milestones: create, update, delete."), "Write milestones: create, update, delete.",
mcp.WithToolAnnotation(annotation.Destructive("Create, update, or delete milestones")), annotation.Destructive("Create, update, or delete milestones"),
mcp.WithString("method", mcp.Required(), mcp.Enum("create", "update", "edit", "delete")), tool.String("method", tool.Required(), tool.Enum("create", "update", "edit", "delete")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("id", mcp.Description("for 'update'/'delete'")), tool.Number("id", tool.Description("for 'update'/'delete'")),
mcp.WithString("title", mcp.Description("for 'create'")), tool.String("title", tool.Description("for 'create'")),
mcp.WithString("description"), tool.String("description"),
mcp.WithString("due_on", mcp.Description("due date")), tool.String("due_on", tool.Description("due date")),
mcp.WithString("state", mcp.Enum("open", "closed")), tool.String("state", tool.Enum("open", "closed")),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: MilestoneReadTool, Tool: MilestoneReadTool,
Handler: milestoneReadFn, Handler: milestoneReadFn,
}) })
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: MilestoneWriteTool, Tool: MilestoneWriteTool,
Handler: milestoneWriteFn, Handler: milestoneWriteFn,
}) })
} }
func milestoneReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func milestoneReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "get": case "get":
return getMilestoneFn(ctx, req) return getMilestoneFn(ctx, args)
case "list": case "list":
return listMilestonesFn(ctx, req) return listMilestonesFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func milestoneWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func milestoneWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "create": case "create":
return createMilestoneFn(ctx, req) return createMilestoneFn(ctx, args)
case "update": case "update":
return editMilestoneFn(ctx, req) return editMilestoneFn(ctx, args)
case "edit": case "edit":
return editMilestoneFn(ctx, req) return editMilestoneFn(ctx, args)
case "delete": case "delete":
return deleteMilestoneFn(ctx, req) return deleteMilestoneFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func getMilestoneFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getMilestoneFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -122,18 +121,18 @@ func getMilestoneFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTool
return to.TextResult(slimMilestone(milestone)) return to.TextResult(slimMilestone(milestone))
} }
func listMilestonesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listMilestonesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
state := params.GetOptionalString(req.GetArguments(), "state", "all") state := params.GetOptionalString(args, "state", "all")
name := params.GetOptionalString(req.GetArguments(), "name", "") name := params.GetOptionalString(args, "name", "")
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
opt := gitea_sdk.ListMilestoneOption{ opt := gitea_sdk.ListMilestoneOption{
State: gitea_sdk.StateType(state), State: gitea_sdk.StateType(state),
Name: name, Name: name,
@@ -153,16 +152,16 @@ func listMilestonesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(slimMilestones(milestones)) return to.TextResult(slimMilestones(milestones))
} }
func createMilestoneFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createMilestoneFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
title, err := params.GetString(req.GetArguments(), "title") title, err := params.GetString(args, "title")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -171,11 +170,11 @@ func createMilestoneFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
Title: title, Title: title,
} }
description, ok := req.GetArguments()["description"].(string) description, ok := args["description"].(string)
if ok { if ok {
opt.Description = description opt.Description = description
} }
opt.Deadline = params.GetOptionalTime(req.GetArguments(), "due_on") opt.Deadline = params.GetOptionalTime(args, "due_on")
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
@@ -189,21 +188,20 @@ func createMilestoneFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
return to.TextResult(slimMilestone(milestone)) return to.TextResult(slimMilestone(milestone))
} }
func editMilestoneFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func editMilestoneFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
args := req.GetArguments()
opt := gitea_sdk.EditMilestoneOption{ opt := gitea_sdk.EditMilestoneOption{
Description: params.GetPresentStringPtr(args, "description"), Description: params.GetPresentStringPtr(args, "description"),
Deadline: params.GetOptionalTime(args, "due_on"), Deadline: params.GetOptionalTime(args, "due_on"),
@@ -228,16 +226,16 @@ func editMilestoneFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToo
return to.TextResult(slimMilestone(milestone)) return to.TextResult(slimMilestone(milestone))
} }
func deleteMilestoneFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteMilestoneFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
+3 -3
View File
@@ -12,7 +12,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
func Test_milestoneWriteFn_dueOn(t *testing.T) { func Test_milestoneWriteFn_dueOn(t *testing.T) {
@@ -56,7 +56,7 @@ func Test_milestoneWriteFn_dueOn(t *testing.T) {
cases := []struct { cases := []struct {
name string name string
fn func(context.Context, mcp.CallToolRequest) (*mcp.CallToolResult, error) fn func(context.Context, map[string]any) (*mcp.CallToolResult, error)
method string method string
extra map[string]any extra map[string]any
}{ }{
@@ -69,7 +69,7 @@ func Test_milestoneWriteFn_dueOn(t *testing.T) {
a := map[string]any{} a := map[string]any{}
maps.Copy(a, args) maps.Copy(a, args)
maps.Copy(a, tc.extra) maps.Copy(a, tc.extra)
res, err := tc.fn(context.Background(), mcp.CallToolRequest{Params: mcp.CallToolParams{Arguments: a}}) res, err := tc.fn(context.Background(), a)
if err != nil || res.IsError { if err != nil || res.IsError {
t.Fatalf("%s err=%v result=%v", tc.name, err, res) t.Fatalf("%s err=%v result=%v", tc.name, err, res)
} }
+36 -41
View File
@@ -12,8 +12,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("notification") var Tool = tool.New("notification")
@@ -24,79 +23,76 @@ const (
) )
var ( var (
NotificationReadTool = mcp.NewTool( NotificationReadTool = tool.NewDefinition(
NotificationReadToolName, NotificationReadToolName,
mcp.WithDescription("Read notifications: list (optionally scoped to a repo) or get a thread by ID."), "Read notifications: list (optionally scoped to a repo) or get a thread by ID.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read notifications")), annotation.ReadOnly("Read notifications"),
mcp.WithString("method", mcp.Required(), mcp.Enum("list", "get")), tool.String("method", tool.Required(), tool.Enum("list", "get")),
mcp.WithString("owner", mcp.Description("scope 'list' to a repo")), tool.String("owner", tool.Description("scope 'list' to a repo")),
mcp.WithString("repo", mcp.Description("scope 'list' to a repo")), tool.String("repo", tool.Description("scope 'list' to a repo")),
mcp.WithNumber("id", mcp.Description("thread ID (for 'get')")), tool.Number("id", tool.Description("thread ID (for 'get')")),
mcp.WithString("status", mcp.Enum("unread", "read", "pinned")), tool.String("status", tool.Enum("unread", "read", "pinned")),
mcp.WithString("subject_type", mcp.Enum("Issue", "Pull", "Commit", "Repository")), tool.String("subject_type", tool.Enum("Issue", "Pull", "Commit", "Repository")),
mcp.WithString("since", mcp.Description("updated after ISO 8601")), tool.String("since", tool.Description("updated after ISO 8601")),
mcp.WithString("before", mcp.Description("updated before ISO 8601")), tool.String("before", tool.Description("updated before ISO 8601")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
NotificationWriteTool = mcp.NewTool( NotificationWriteTool = tool.NewDefinition(
NotificationWriteToolName, NotificationWriteToolName,
mcp.WithDescription("Mark a notification or all notifications as read."), "Mark a notification or all notifications as read.",
mcp.WithToolAnnotation(annotation.Write("Manage notifications")), annotation.Write("Manage notifications"),
mcp.WithString("method", mcp.Required(), mcp.Enum("mark_read", "mark_all_read")), tool.String("method", tool.Required(), tool.Enum("mark_read", "mark_all_read")),
mcp.WithNumber("id", mcp.Description("thread ID (for 'mark_read')")), tool.Number("id", tool.Description("thread ID (for 'mark_read')")),
mcp.WithString("owner", mcp.Description("scope 'mark_all_read' to a repo")), tool.String("owner", tool.Description("scope 'mark_all_read' to a repo")),
mcp.WithString("repo", mcp.Description("scope 'mark_all_read' to a repo")), tool.String("repo", tool.Description("scope 'mark_all_read' to a repo")),
mcp.WithString("last_read_at", mcp.Description("ISO 8601; defaults to now")), tool.String("last_read_at", tool.Description("ISO 8601; defaults to now")),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: NotificationReadTool, Tool: NotificationReadTool,
Handler: notificationReadFn, Handler: notificationReadFn,
}) })
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: NotificationWriteTool, Tool: NotificationWriteTool,
Handler: notificationWriteFn, Handler: notificationWriteFn,
}) })
} }
func notificationReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func notificationReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
method, err := params.GetString(args, "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "list": case "list":
return listNotificationsFn(ctx, req) return listNotificationsFn(ctx, args)
case "get": case "get":
return getNotificationFn(ctx, req) return getNotificationFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func notificationWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func notificationWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
method, err := params.GetString(args, "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "mark_read": case "mark_read":
return markNotificationReadFn(ctx, req) return markNotificationReadFn(ctx, args)
case "mark_all_read": case "mark_all_read":
return markAllNotificationsReadFn(ctx, req) return markAllNotificationsReadFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func listNotificationsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listNotificationsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
page, pageSize := params.GetPagination(args, 30) page, pageSize := params.GetPagination(args, 30)
opt := gitea_sdk.ListNotificationOptions{ opt := gitea_sdk.ListNotificationOptions{
ListOptions: gitea_sdk.ListOptions{ ListOptions: gitea_sdk.ListOptions{
@@ -139,8 +135,8 @@ func listNotificationsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Cal
return to.TextResult(slimThreads(threads)) return to.TextResult(slimThreads(threads))
} }
func getNotificationFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getNotificationFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -155,8 +151,8 @@ func getNotificationFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
return to.TextResult(slimThread(thread)) return to.TextResult(slimThread(thread))
} }
func markNotificationReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func markNotificationReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -174,8 +170,7 @@ func markNotificationReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.
return to.TextResult("Notification marked as read") return to.TextResult("Notification marked as read")
} }
func markAllNotificationsReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func markAllNotificationsReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
lastReadAt := time.Now() lastReadAt := time.Now()
if t := params.GetOptionalTime(args, "last_read_at"); t != nil { if t := params.GetOptionalTime(args, "last_read_at"); t != nil {
lastReadAt = *t lastReadAt = *t
+60 -23
View File
@@ -29,11 +29,23 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/log" "gitea.com/gitea/gitea-mcp/pkg/log"
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
"github.com/mark3labs/mcp-go/server" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
// maxRequestBodyBytes raises the SDK's 4 MiB default, which is too tight for the
// base64 file content create_or_update_file accepts.
const maxRequestBodyBytes = 32 << 20
// sessionTimeout expires idle sessions, which the SDK otherwise keeps for the
// process lifetime: a client that goes away without DELETE /mcp leaks its
// session, and initialize takes no token. Clients re-initialize on the 404.
const sessionTimeout = 30 * time.Minute
// httpReadHeaderTimeout bounds slow header reads without limiting SSE writes.
const httpReadHeaderTimeout = 10 * time.Second
var ( var (
mcpServer *server.MCPServer mcpServer *mcp.Server
domainTools = []*tool.Tool{ domainTools = []*tool.Tool{
user.Tool, actions.Tool, repo.Tool, notification.Tool, issue.Tool, user.Tool, actions.Tool, repo.Tool, notification.Tool, issue.Tool,
@@ -43,9 +55,11 @@ var (
} }
) )
func RegisterTool(s *server.MCPServer) { func RegisterTool(s *mcp.Server) {
for _, t := range domainTools { for _, t := range domainTools {
s.AddTools(t.Tools()...) for _, registeredTool := range t.Tools() {
s.AddTool(registeredTool.Tool, registeredTool.MCPHandler())
}
} }
tool.WarnUnmatchedAllowedTools(domainTools...) tool.WarnUnmatchedAllowedTools(domainTools...)
tool.WarnUnmatchedAllowedScopes(domainTools...) tool.WarnUnmatchedAllowedScopes(domainTools...)
@@ -71,8 +85,7 @@ func parseAuthToken(authHeader string) (string, bool) {
return "", false return "", false
} }
func getContextWithToken(ctx context.Context, r *http.Request) context.Context { func getContextWithToken(ctx context.Context, authHeader string) context.Context {
authHeader := r.Header.Get("Authorization")
if authHeader == "" { if authHeader == "" {
return ctx return ctx
} }
@@ -85,23 +98,43 @@ func getContextWithToken(ctx context.Context, r *http.Request) context.Context {
return context.WithValue(ctx, mcpContext.TokenContextKey, token) return context.WithValue(ctx, mcpContext.TokenContextKey, token)
} }
func authTokenMiddleware(next mcp.MethodHandler) mcp.MethodHandler {
return func(ctx context.Context, method string, req mcp.Request) (mcp.Result, error) {
if extra := req.GetExtra(); extra != nil {
ctx = getContextWithToken(ctx, extra.Header.Get("Authorization"))
}
return next(ctx, method, req)
}
}
func newHTTPServer(addr string, s *mcp.Server) *http.Server {
mux := http.NewServeMux()
mux.Handle("/mcp", mcp.NewStreamableHTTPHandler(
func(*http.Request) *mcp.Server { return s },
&mcp.StreamableHTTPOptions{
Logger: log.Slog(),
MaxRequestBodyBytes: maxRequestBodyBytes,
Stateless: false, // SessionTimeout requires stateful sessions.
SessionTimeout: sessionTimeout,
},
))
return &http.Server{
Addr: addr,
Handler: mux,
ReadHeaderTimeout: httpReadHeaderTimeout,
}
}
func Run() error { func Run() error {
mcpServer = newMCPServer(flag.Version) mcpServer = newMCPServer(flag.Version)
RegisterTool(mcpServer) RegisterTool(mcpServer)
switch flag.Mode { switch flag.Mode {
case "stdio": case "stdio":
if err := server.ServeStdio( if err := mcpServer.Run(context.Background(), &mcp.StdioTransport{}); err != nil {
mcpServer,
); err != nil {
return err return err
} }
case "http": case "http":
httpServer := server.NewStreamableHTTPServer( httpServer := newHTTPServer(fmt.Sprintf(":%d", flag.Port), mcpServer)
mcpServer,
server.WithStreamableHTTPLogger(log.Slog()),
server.WithHeartbeatInterval(30*time.Second),
server.WithHTTPContextFunc(getContextWithToken),
)
log.Infof("Gitea MCP HTTP server listening on :%d", flag.Port) log.Infof("Gitea MCP HTTP server listening on :%d", flag.Port)
// Graceful shutdown setup // Graceful shutdown setup
@@ -120,7 +153,7 @@ func Run() error {
close(shutdownDone) close(shutdownDone)
}() }()
if err := httpServer.Start(fmt.Sprintf(":%d", flag.Port)); err != nil && !errors.Is(err, http.ErrServerClosed) { if err := httpServer.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
return err return err
} }
<-shutdownDone // Wait for shutdown to finish <-shutdownDone // Wait for shutdown to finish
@@ -130,12 +163,16 @@ func Run() error {
return nil return nil
} }
func newMCPServer(version string) *server.MCPServer { func newMCPServer(version string) *mcp.Server {
return server.NewMCPServer( // SDK keepalives send MCP ping requests and disconnect clients without a
"Gitea MCP Server", // server-to-client channel, so KeepAlive stays disabled.
version, s := mcp.NewServer(
server.WithToolCapabilities(true), &mcp.Implementation{
server.WithLogging(), Name: "Gitea MCP Server",
server.WithRecovery(), Version: version,
},
&mcp.ServerOptions{Logger: log.Slog()},
) )
s.AddReceivingMiddleware(authTokenMiddleware)
return s
} }
+12 -46
View File
@@ -1,54 +1,20 @@
package operation package operation
import ( import "testing"
"testing"
"gitea.com/gitea/gitea-mcp/pkg/flag" func TestNewHTTPServerConfig(t *testing.T) {
) server := newHTTPServer(":12345", newMCPServer("test"))
if server.Addr != ":12345" {
// TestAllToolsHaveDescriptions ensures every registered tool sets a non-empty t.Errorf("Addr = %q, want %q", server.Addr, ":12345")
// Tool.Description. mcp-go only serializes the "description" field of a tool
// when it is non-empty, so an omitted description makes strict MCP clients
// (e.g. mcp-probe) reject the tools/list response with "missing field
// `description`".
func TestAllToolsHaveDescriptions(t *testing.T) {
origRO, origAllow := flag.ReadOnly, flag.AllowedTools
t.Cleanup(func() {
flag.ReadOnly, flag.AllowedTools = origRO, origAllow
})
flag.ReadOnly = false
flag.AllowedTools = nil
var missing []string
for _, d := range domainTools {
for _, st := range d.Tools() {
if st.Tool.Description == "" {
missing = append(missing, st.Tool.Name)
}
}
} }
if len(missing) > 0 { if server.Handler == nil {
t.Errorf("tools missing a description: %v", missing) t.Error("Handler is nil")
} }
} if server.ReadHeaderTimeout != httpReadHeaderTimeout {
t.Errorf("ReadHeaderTimeout = %v, want %v", server.ReadHeaderTimeout, httpReadHeaderTimeout)
// TestDomainToolsScopesAreUniqueAndNonEmpty ensures every entry registered in }
// domainTools has a canonical, non-empty scope name and that no two domains if server.WriteTimeout != 0 {
// share the same scope (each domain.Tools() call is filtered by exactly one t.Errorf("WriteTimeout = %v, want zero for SSE", server.WriteTimeout)
// scope name via flag.AllowedScopes).
func TestDomainToolsScopesAreUniqueAndNonEmpty(t *testing.T) {
seen := map[string]struct{}{}
for _, d := range domainTools {
scope := d.Scope()
if scope == "" {
t.Errorf("domainTools contains a domain with an empty scope")
continue
}
if _, ok := seen[scope]; ok {
t.Errorf("domainTools contains a duplicate scope %q", scope)
continue
}
seen[scope] = struct{}{}
} }
} }
+32 -39
View File
@@ -13,8 +13,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/to" "gitea.com/gitea/gitea-mcp/pkg/to"
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("packages") var Tool = tool.New("packages")
@@ -25,70 +24,68 @@ const (
) )
var ( var (
PackageReadTool = mcp.NewTool( PackageReadTool = tool.NewDefinition(
PackageReadToolName, PackageReadToolName,
mcp.WithToolAnnotation(annotation.ReadOnly("Read package registry")), "Read package registry: list packages (one entry per version, filter via 'q'/'type'), list versions, or get a version.",
mcp.WithDescription("Read package registry: list packages (one entry per version, filter via 'q'/'type'), list versions, or get a version."), annotation.ReadOnly("Read package registry"),
mcp.WithString("method", mcp.Required(), mcp.Enum("list", "list_versions", "get")), tool.String("method", tool.Required(), tool.Enum("list", "list_versions", "get")),
mcp.WithString("owner", mcp.Required(), mcp.Description("user or org")), tool.String("owner", tool.Required(), tool.Description("user or org")),
mcp.WithString("type", mcp.Description("container/npm/maven/pypi/cargo/generic; required except 'list'")), tool.String("type", tool.Description("container/npm/maven/pypi/cargo/generic; required except 'list'")),
mcp.WithString("name", mcp.Description("slashes auto-encoded; required except 'list'")), tool.String("name", tool.Description("slashes auto-encoded; required except 'list'")),
mcp.WithString("version", mcp.Description("for 'get'")), tool.String("version", tool.Description("for 'get'")),
mcp.WithString("q", mcp.Description("search query")), tool.String("q", tool.Description("search query")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1), mcp.Min(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1), tool.Minimum(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30), mcp.Min(1)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30), tool.Minimum(1)),
) )
PackageWriteTool = mcp.NewTool( PackageWriteTool = tool.NewDefinition(
PackageWriteToolName, PackageWriteToolName,
mcp.WithToolAnnotation(annotation.Destructive("Delete a package version")), "Delete a package version (irreversible).",
mcp.WithDescription("Delete a package version (irreversible)."), annotation.Destructive("Delete a package version"),
mcp.WithString("method", mcp.Required(), mcp.Enum("delete")), tool.String("method", tool.Required(), tool.Enum("delete")),
mcp.WithString("owner", mcp.Required(), mcp.Description("user or org")), tool.String("owner", tool.Required(), tool.Description("user or org")),
mcp.WithString("type", mcp.Required(), mcp.Description("container/npm/maven/pypi/cargo/generic")), tool.String("type", tool.Required(), tool.Description("container/npm/maven/pypi/cargo/generic")),
mcp.WithString("name", mcp.Required(), mcp.Description("slashes auto-encoded")), tool.String("name", tool.Required(), tool.Description("slashes auto-encoded")),
mcp.WithString("version", mcp.Required()), tool.String("version", tool.Required()),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: PackageReadTool, Tool: PackageReadTool,
Handler: packageReadFn, Handler: packageReadFn,
}) })
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: PackageWriteTool, Tool: PackageWriteTool,
Handler: packageWriteFn, Handler: packageWriteFn,
}) })
} }
func packageReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func packageReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
method, err := params.GetString(args, "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "list": case "list":
return listPackagesFn(ctx, req) return listPackagesFn(ctx, args)
case "list_versions": case "list_versions":
return listPackageVersionsFn(ctx, req) return listPackageVersionsFn(ctx, args)
case "get": case "get":
return getPackageFn(ctx, req) return getPackageFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func packageWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func packageWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
method, err := params.GetString(args, "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "delete": case "delete":
return deletePackageVersionFn(ctx, req) return deletePackageVersionFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
@@ -108,8 +105,7 @@ func escapePackageName(name string) string {
return url.PathEscape(name) return url.PathEscape(name)
} }
func listPackagesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listPackagesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -135,8 +131,7 @@ func listPackagesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTool
return to.TextResult(slimPackages(result)) return to.TextResult(slimPackages(result))
} }
func listPackageVersionsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listPackageVersionsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -164,8 +159,7 @@ func listPackageVersionsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.C
return to.TextResult(slimPackages(result)) return to.TextResult(slimPackages(result))
} }
func getPackageFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getPackageFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -192,8 +186,7 @@ func getPackageFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolRe
return to.TextResult(slimPackage(result)) return to.TextResult(slimPackage(result))
} }
func deletePackageVersionFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deletePackageVersionFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+20 -28
View File
@@ -11,7 +11,7 @@ import (
mcpContext "gitea.com/gitea/gitea-mcp/pkg/context" mcpContext "gitea.com/gitea/gitea-mcp/pkg/context"
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
func TestPackageReadList(t *testing.T) { func TestPackageReadList(t *testing.T) {
@@ -37,13 +37,12 @@ func TestPackageReadList(t *testing.T) {
ctx := context.WithValue(context.Background(), mcpContext.TokenContextKey, "test-token") ctx := context.WithValue(context.Background(), mcpContext.TokenContextKey, "test-token")
t.Run("basic list", func(t *testing.T) { t.Run("basic list", func(t *testing.T) {
req := mcp.CallToolRequest{} args := map[string]any{
req.Params.Arguments = map[string]any{
"method": "list", "method": "list",
"owner": "test-org", "owner": "test-org",
} }
result, err := packageReadFn(ctx, req) result, err := packageReadFn(ctx, args)
if err != nil { if err != nil {
t.Fatalf("packageReadFn() error: %v", err) t.Fatalf("packageReadFn() error: %v", err)
} }
@@ -51,7 +50,7 @@ func TestPackageReadList(t *testing.T) {
t.Fatal("packageReadFn() returned error result") t.Fatal("packageReadFn() returned error result")
} }
text := result.Content[0].(mcp.TextContent).Text text := result.Content[0].(*mcp.TextContent).Text
var packages []map[string]any var packages []map[string]any
if err := json.Unmarshal([]byte(text), &packages); err != nil { if err := json.Unmarshal([]byte(text), &packages); err != nil {
t.Fatalf("failed to unmarshal result: %v", err) t.Fatalf("failed to unmarshal result: %v", err)
@@ -68,15 +67,14 @@ func TestPackageReadList(t *testing.T) {
}) })
t.Run("with type and query filters", func(t *testing.T) { t.Run("with type and query filters", func(t *testing.T) {
req := mcp.CallToolRequest{} args := map[string]any{
req.Params.Arguments = map[string]any{
"method": "list", "method": "list",
"owner": "test-org", "owner": "test-org",
"type": "container", "type": "container",
"q": "myimage", "q": "myimage",
} }
_, err := packageReadFn(ctx, req) _, err := packageReadFn(ctx, args)
if err != nil { if err != nil {
t.Fatalf("packageReadFn() error: %v", err) t.Fatalf("packageReadFn() error: %v", err)
} }
@@ -92,15 +90,14 @@ func TestPackageReadList(t *testing.T) {
}) })
t.Run("with pagination", func(t *testing.T) { t.Run("with pagination", func(t *testing.T) {
req := mcp.CallToolRequest{} args := map[string]any{
req.Params.Arguments = map[string]any{
"method": "list", "method": "list",
"owner": "test-org", "owner": "test-org",
"page": float64(2), "page": float64(2),
"per_page": float64(10), "per_page": float64(10),
} }
_, err := packageReadFn(ctx, req) _, err := packageReadFn(ctx, args)
if err != nil { if err != nil {
t.Fatalf("packageReadFn() error: %v", err) t.Fatalf("packageReadFn() error: %v", err)
} }
@@ -148,15 +145,14 @@ func TestPackageReadListVersions(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.testName, func(t *testing.T) { t.Run(tt.testName, func(t *testing.T) {
req := mcp.CallToolRequest{} args := map[string]any{
req.Params.Arguments = map[string]any{
"method": "list_versions", "method": "list_versions",
"owner": "test-org", "owner": "test-org",
"type": "container", "type": "container",
"name": tt.name, "name": tt.name,
} }
result, err := packageReadFn(ctx, req) result, err := packageReadFn(ctx, args)
if err != nil { if err != nil {
t.Fatalf("packageReadFn() error: %v", err) t.Fatalf("packageReadFn() error: %v", err)
} }
@@ -171,7 +167,7 @@ func TestPackageReadListVersions(t *testing.T) {
} }
mu.Unlock() mu.Unlock()
text := result.Content[0].(mcp.TextContent).Text text := result.Content[0].(*mcp.TextContent).Text
var versions []map[string]any var versions []map[string]any
if err := json.Unmarshal([]byte(text), &versions); err != nil { if err := json.Unmarshal([]byte(text), &versions); err != nil {
t.Fatalf("failed to unmarshal result: %v", err) t.Fatalf("failed to unmarshal result: %v", err)
@@ -215,8 +211,7 @@ func TestPackageReadGet(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.testName, func(t *testing.T) { t.Run(tt.testName, func(t *testing.T) {
req := mcp.CallToolRequest{} args := map[string]any{
req.Params.Arguments = map[string]any{
"method": "get", "method": "get",
"owner": "test-org", "owner": "test-org",
"type": "container", "type": "container",
@@ -224,7 +219,7 @@ func TestPackageReadGet(t *testing.T) {
"version": "v1.0.0", "version": "v1.0.0",
} }
result, err := packageReadFn(ctx, req) result, err := packageReadFn(ctx, args)
if err != nil { if err != nil {
t.Fatalf("packageReadFn() error: %v", err) t.Fatalf("packageReadFn() error: %v", err)
} }
@@ -239,7 +234,7 @@ func TestPackageReadGet(t *testing.T) {
} }
mu.Unlock() mu.Unlock()
text := result.Content[0].(mcp.TextContent).Text text := result.Content[0].(*mcp.TextContent).Text
var pkg map[string]any var pkg map[string]any
if err := json.Unmarshal([]byte(text), &pkg); err != nil { if err := json.Unmarshal([]byte(text), &pkg); err != nil {
t.Fatalf("failed to unmarshal result: %v", err) t.Fatalf("failed to unmarshal result: %v", err)
@@ -277,8 +272,7 @@ func TestPackageWriteDelete(t *testing.T) {
ctx := context.WithValue(context.Background(), mcpContext.TokenContextKey, "test-token") ctx := context.WithValue(context.Background(), mcpContext.TokenContextKey, "test-token")
req := mcp.CallToolRequest{} args := map[string]any{
req.Params.Arguments = map[string]any{
"method": "delete", "method": "delete",
"owner": "test-org", "owner": "test-org",
"type": "container", "type": "container",
@@ -286,7 +280,7 @@ func TestPackageWriteDelete(t *testing.T) {
"version": "v1.0.0", "version": "v1.0.0",
} }
result, err := packageWriteFn(ctx, req) result, err := packageWriteFn(ctx, args)
if err != nil { if err != nil {
t.Fatalf("packageWriteFn() error: %v", err) t.Fatalf("packageWriteFn() error: %v", err)
} }
@@ -307,24 +301,22 @@ func TestPackageWriteDelete(t *testing.T) {
func TestPackageReadUnknownMethod(t *testing.T) { func TestPackageReadUnknownMethod(t *testing.T) {
ctx := context.Background() ctx := context.Background()
req := mcp.CallToolRequest{} args := map[string]any{
req.Params.Arguments = map[string]any{
"method": "bogus", "method": "bogus",
"owner": "test-org", "owner": "test-org",
} }
if _, err := packageReadFn(ctx, req); err == nil { if _, err := packageReadFn(ctx, args); err == nil {
t.Fatal("expected error for unknown method") t.Fatal("expected error for unknown method")
} }
} }
func TestPackageWriteUnknownMethod(t *testing.T) { func TestPackageWriteUnknownMethod(t *testing.T) {
ctx := context.Background() ctx := context.Background()
req := mcp.CallToolRequest{} args := map[string]any{
req.Params.Arguments = map[string]any{
"method": "bogus", "method": "bogus",
"owner": "test-org", "owner": "test-org",
} }
if _, err := packageWriteFn(ctx, req); err == nil { if _, err := packageWriteFn(ctx, args); err == nil {
t.Fatal("expected error for unknown method") t.Fatal("expected error for unknown method")
} }
} }
+131 -151
View File
@@ -15,8 +15,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("pull_request") var Tool = tool.New("pull_request")
@@ -29,79 +28,79 @@ const (
) )
var ( var (
ListRepoPullRequestsTool = mcp.NewTool( ListRepoPullRequestsTool = tool.NewDefinition(
ListRepoPullRequestsToolName, ListRepoPullRequestsToolName,
mcp.WithDescription("List pull requests in a repository, filterable by state and milestone, with configurable sort order (e.g. recently updated, most commented)."), "List pull requests in a repository, filterable by state and milestone, with configurable sort order (e.g. recently updated, most commented).",
mcp.WithToolAnnotation(annotation.ReadOnly("List pull requests")), annotation.ReadOnly("List pull requests"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("state", mcp.Enum("open", "closed", "all"), mcp.DefaultString("all")), tool.String("state", tool.Enum("open", "closed", "all"), tool.Default("all")),
mcp.WithString("sort", mcp.Enum("oldest", "recentupdate", "leastupdate", "mostcomment", "leastcomment", "priority"), mcp.DefaultString("recentupdate")), tool.String("sort", tool.Enum("oldest", "recentupdate", "leastupdate", "mostcomment", "leastcomment", "priority"), tool.Default("recentupdate")),
mcp.WithNumber("milestone"), tool.Number("milestone"),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
PullRequestReadTool = mcp.NewTool( PullRequestReadTool = tool.NewDefinition(
PullRequestReadToolName, PullRequestReadToolName,
mcp.WithDescription("Read pull request: details, diff, changed files, head commit status, reviews, review comments."), "Read pull request: details, diff, changed files, head commit status, reviews, review comments.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read pull request details")), annotation.ReadOnly("Read pull request details"),
mcp.WithString("method", mcp.Required(), mcp.Enum("get", "get_diff", "get_files", "get_status", "get_reviews", "get_review", "get_review_comments")), tool.String("method", tool.Required(), tool.Enum("get", "get_diff", "get_files", "get_status", "get_reviews", "get_review", "get_review_comments")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("pull_number", mcp.Required()), tool.Number("pull_number", tool.Required()),
mcp.WithNumber("review_id", mcp.Description("for 'get_review'; optional for 'get_review_comments', omit to list all")), tool.Number("review_id", tool.Description("for 'get_review'; optional for 'get_review_comments', omit to list all")),
mcp.WithBoolean("binary", mcp.Description("include binary diff")), tool.Boolean("binary", tool.Description("include binary diff")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
PullRequestWriteTool = mcp.NewTool( PullRequestWriteTool = tool.NewDefinition(
PullRequestWriteToolName, PullRequestWriteToolName,
mcp.WithDescription("Write pull requests: create, update, close, reopen, merge, update branch from base, manage reviewers."), "Write pull requests: create, update, close, reopen, merge, update branch from base, manage reviewers.",
mcp.WithToolAnnotation(annotation.Write("Create, update, close, reopen, or merge pull requests")), annotation.Write("Create, update, close, reopen, or merge pull requests"),
mcp.WithString("method", mcp.Required(), mcp.Enum("create", "update", "close", "reopen", "merge", "update_branch", "add_reviewers", "remove_reviewers")), tool.String("method", tool.Required(), tool.Enum("create", "update", "close", "reopen", "merge", "update_branch", "add_reviewers", "remove_reviewers")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("pull_number", mcp.Description("required except for 'create'")), tool.Number("pull_number", tool.Description("required except for 'create'")),
mcp.WithString("title", mcp.Description("required for 'create'; optional for 'update'/'merge'")), tool.String("title", tool.Description("required for 'create'; optional for 'update'/'merge'")),
mcp.WithString("body", mcp.Description("required for 'create'; optional for 'update'")), tool.String("body", tool.Description("required for 'create'; optional for 'update'")),
mcp.WithString("head", mcp.Description("head branch (required for 'create')")), tool.String("head", tool.Description("head branch (required for 'create')")),
mcp.WithString("base", mcp.Description("base branch (required for 'create')")), tool.String("base", tool.Description("base branch (required for 'create')")),
mcp.WithString("assignee", mcp.Description("for 'update'")), tool.String("assignee", tool.Description("for 'update'")),
mcp.WithArray("assignees", mcp.Description("for 'update'"), mcp.Items(map[string]any{"type": "string"})), tool.Array("assignees", tool.Description("for 'update'"), tool.Items(map[string]any{"type": "string"})),
mcp.WithNumber("milestone", mcp.Description("for 'update'")), tool.Number("milestone", tool.Description("for 'update'")),
mcp.WithString("state", mcp.Description("for 'update'"), mcp.Enum("open", "closed")), tool.String("state", tool.Description("for 'update'"), tool.Enum("open", "closed")),
mcp.WithBoolean("allow_maintainer_edit", mcp.Description("for 'update'")), tool.Boolean("allow_maintainer_edit", tool.Description("for 'update'")),
mcp.WithArray("labels", mcp.Description("label IDs"), mcp.Items(map[string]any{"type": "number"})), tool.Array("labels", tool.Description("label IDs"), tool.Items(map[string]any{"type": "number"})),
mcp.WithString("deadline", mcp.Description("ISO 8601")), tool.String("deadline", tool.Description("ISO 8601")),
mcp.WithBoolean("remove_deadline", mcp.Description("for 'update'")), tool.Boolean("remove_deadline", tool.Description("for 'update'")),
mcp.WithString("merge_style", mcp.Description("for 'merge'"), mcp.Enum("merge", "rebase", "rebase-merge", "squash", "fast-forward-only"), mcp.DefaultString("merge")), tool.String("merge_style", tool.Description("for 'merge'"), tool.Enum("merge", "rebase", "rebase-merge", "squash", "fast-forward-only"), tool.Default("merge")),
mcp.WithString("message", mcp.Description("merge commit message or dismissal reason")), tool.String("message", tool.Description("merge commit message or dismissal reason")),
mcp.WithBoolean("delete_branch", mcp.Description("for 'merge'")), tool.Boolean("delete_branch", tool.Description("for 'merge'")),
mcp.WithBoolean("force_merge", mcp.Description("merge even if checks fail")), tool.Boolean("force_merge", tool.Description("merge even if checks fail")),
mcp.WithBoolean("merge_when_checks_succeed", mcp.Description("for 'merge'")), tool.Boolean("merge_when_checks_succeed", tool.Description("for 'merge'")),
mcp.WithString("head_commit_id", mcp.Description("expected head SHA for conflict detection")), tool.String("head_commit_id", tool.Description("expected head SHA for conflict detection")),
mcp.WithArray("reviewers", mcp.Description("for 'add_reviewers'/'remove_reviewers'"), mcp.Items(map[string]any{"type": "string"})), tool.Array("reviewers", tool.Description("for 'add_reviewers'/'remove_reviewers'"), tool.Items(map[string]any{"type": "string"})),
mcp.WithArray("team_reviewers", mcp.Description("for 'add_reviewers'/'remove_reviewers'"), mcp.Items(map[string]any{"type": "string"})), tool.Array("team_reviewers", tool.Description("for 'add_reviewers'/'remove_reviewers'"), tool.Items(map[string]any{"type": "string"})),
mcp.WithBoolean("draft", mcp.Description("uses 'WIP: ' title prefix")), tool.Boolean("draft", tool.Description("uses 'WIP: ' title prefix")),
) )
PullRequestReviewWriteTool = mcp.NewTool( PullRequestReviewWriteTool = tool.NewDefinition(
PullRequestReviewWriteToolName, PullRequestReviewWriteToolName,
mcp.WithDescription("Write PR reviews: create, submit, delete, dismiss, reply to and resolve review comments."), "Write PR reviews: create, submit, delete, dismiss, reply to and resolve review comments.",
mcp.WithToolAnnotation(annotation.Write("Write pull request reviews")), annotation.Write("Write pull request reviews"),
mcp.WithString("method", mcp.Required(), mcp.Enum("create", "submit", "delete", "dismiss", "reply_comment", "resolve_thread", "unresolve_thread")), tool.String("method", tool.Required(), tool.Enum("create", "submit", "delete", "dismiss", "reply_comment", "resolve_thread", "unresolve_thread")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("pull_number", mcp.Description("required except for 'resolve_thread'/'unresolve_thread'")), tool.Number("pull_number", tool.Description("required except for 'resolve_thread'/'unresolve_thread'")),
mcp.WithNumber("review_id", mcp.Description("for 'submit'/'delete'/'dismiss'")), tool.Number("review_id", tool.Description("for 'submit'/'delete'/'dismiss'")),
mcp.WithNumber("comment_id", mcp.Description("comment ID from 'get_review_comments'; resolve takes the thread's first")), tool.Number("comment_id", tool.Description("comment ID from 'get_review_comments'; resolve takes the thread's first")),
mcp.WithString("state", mcp.Enum("APPROVED", "REQUEST_CHANGES", "COMMENT", "PENDING")), tool.String("state", tool.Enum("APPROVED", "REQUEST_CHANGES", "COMMENT", "PENDING")),
mcp.WithString("body", mcp.Description("review body, or reply text for 'reply_comment'")), tool.String("body", tool.Description("review body, or reply text for 'reply_comment'")),
mcp.WithString("commit_id", mcp.Description("for 'create'")), tool.String("commit_id", tool.Description("for 'create'")),
mcp.WithString("message", mcp.Description("dismissal reason")), tool.String("message", tool.Description("dismissal reason")),
mcp.WithArray("comments", mcp.Description("inline comments (for 'create')"), mcp.Items(map[string]any{ tool.Array("comments", tool.Description("inline comments (for 'create')"), tool.Items(map[string]any{
"type": "object", "type": "object",
"properties": map[string]any{ "properties": map[string]any{
"path": map[string]any{"type": "string"}, "path": map[string]any{"type": "string"},
@@ -114,86 +113,86 @@ var (
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: ListRepoPullRequestsTool, Tool: ListRepoPullRequestsTool,
Handler: listRepoPullRequestsFn, Handler: listRepoPullRequestsFn,
}) })
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: PullRequestReadTool, Tool: PullRequestReadTool,
Handler: pullRequestReadFn, Handler: pullRequestReadFn,
}) })
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: PullRequestWriteTool, Tool: PullRequestWriteTool,
Handler: pullRequestWriteFn, Handler: pullRequestWriteFn,
}) })
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: PullRequestReviewWriteTool, Tool: PullRequestReviewWriteTool,
Handler: pullRequestReviewWriteFn, Handler: pullRequestReviewWriteFn,
}) })
} }
func pullRequestReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func pullRequestReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "get": case "get":
return getPullRequestByIndexFn(ctx, req) return getPullRequestByIndexFn(ctx, args)
case "get_diff": case "get_diff":
return getPullRequestDiffFn(ctx, req) return getPullRequestDiffFn(ctx, args)
case "get_files": case "get_files":
return getPullRequestFilesFn(ctx, req) return getPullRequestFilesFn(ctx, args)
case "get_status": case "get_status":
return getPullRequestStatusFn(ctx, req) return getPullRequestStatusFn(ctx, args)
case "get_reviews": case "get_reviews":
return listPullRequestReviewsFn(ctx, req) return listPullRequestReviewsFn(ctx, args)
case "get_review": case "get_review":
return getPullRequestReviewFn(ctx, req) return getPullRequestReviewFn(ctx, args)
case "get_review_comments": case "get_review_comments":
return listPullRequestReviewCommentsFn(ctx, req) return listPullRequestReviewCommentsFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func pullRequestWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func pullRequestWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "create": case "create":
return createPullRequestFn(ctx, req) return createPullRequestFn(ctx, args)
case "update": case "update":
return editPullRequestFn(ctx, req) return editPullRequestFn(ctx, args)
case "close": case "close":
return closePullRequestFn(ctx, req) return closePullRequestFn(ctx, args)
case "reopen": case "reopen":
return reopenPullRequestFn(ctx, req) return reopenPullRequestFn(ctx, args)
case "merge": case "merge":
return mergePullRequestFn(ctx, req) return mergePullRequestFn(ctx, args)
case "update_branch": case "update_branch":
return updatePullRequestBranchFn(ctx, req) return updatePullRequestBranchFn(ctx, args)
case "add_reviewers": case "add_reviewers":
return createPullRequestReviewerFn(ctx, req) return createPullRequestReviewerFn(ctx, args)
case "remove_reviewers": case "remove_reviewers":
return deletePullRequestReviewerFn(ctx, req) return deletePullRequestReviewerFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func closePullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func closePullRequestFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "pull_number") index, err := params.GetIndex(args, "pull_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -214,16 +213,16 @@ func closePullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Call
return to.TextResult(slimPullRequest(pr)) return to.TextResult(slimPullRequest(pr))
} }
func reopenPullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func reopenPullRequestFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "pull_number") index, err := params.GetIndex(args, "pull_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -244,33 +243,32 @@ func reopenPullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Cal
return to.TextResult(slimPullRequest(pr)) return to.TextResult(slimPullRequest(pr))
} }
func pullRequestReviewWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func pullRequestReviewWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "create": case "create":
return createPullRequestReviewFn(ctx, req) return createPullRequestReviewFn(ctx, args)
case "submit": case "submit":
return submitPullRequestReviewFn(ctx, req) return submitPullRequestReviewFn(ctx, args)
case "delete": case "delete":
return deletePullRequestReviewFn(ctx, req) return deletePullRequestReviewFn(ctx, args)
case "dismiss": case "dismiss":
return dismissPullRequestReviewFn(ctx, req) return dismissPullRequestReviewFn(ctx, args)
case "reply_comment": case "reply_comment":
return replyPullRequestReviewCommentFn(ctx, req) return replyPullRequestReviewCommentFn(ctx, args)
case "resolve_thread": case "resolve_thread":
return resolveReviewThreadFn(ctx, req) return resolveReviewThreadFn(ctx, args)
case "unresolve_thread": case "unresolve_thread":
return unresolveReviewThreadFn(ctx, req) return unresolveReviewThreadFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func getPullRequestByIndexFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getPullRequestByIndexFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -305,8 +303,7 @@ func getPullRequestByIndexFn(ctx context.Context, req mcp.CallToolRequest) (*mcp
return to.TextResult(m) return to.TextResult(m)
} }
func getPullRequestDiffFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getPullRequestDiffFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -335,8 +332,7 @@ func getPullRequestDiffFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Ca
return to.TextResult(string(diffBytes)) return to.TextResult(string(diffBytes))
} }
func listRepoPullRequestsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoPullRequestsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -392,8 +388,7 @@ func applyDraftPrefix(title string, isDraft bool) string {
return title return title
} }
func createPullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createPullRequestFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -447,8 +442,7 @@ func createPullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Cal
type reviewerOp func(client *gitea_sdk.PullRequestsService, ctx context.Context, owner, repo string, index int64, opt gitea_sdk.PullReviewRequestOptions) (*gitea_sdk.Response, error) type reviewerOp func(client *gitea_sdk.PullRequestsService, ctx context.Context, owner, repo string, index int64, opt gitea_sdk.PullReviewRequestOptions) (*gitea_sdk.Response, error)
func pullRequestReviewerFn(ctx context.Context, req mcp.CallToolRequest, verb string, op reviewerOp) (*mcp.CallToolResult, error) { func pullRequestReviewerFn(ctx context.Context, args map[string]any, verb string, op reviewerOp) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -486,16 +480,15 @@ func pullRequestReviewerFn(ctx context.Context, req mcp.CallToolRequest, verb st
}) })
} }
func createPullRequestReviewerFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createPullRequestReviewerFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
return pullRequestReviewerFn(ctx, req, "create", (*gitea_sdk.PullRequestsService).CreateReviewRequests) return pullRequestReviewerFn(ctx, args, "create", (*gitea_sdk.PullRequestsService).CreateReviewRequests)
} }
func deletePullRequestReviewerFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deletePullRequestReviewerFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
return pullRequestReviewerFn(ctx, req, "delete", (*gitea_sdk.PullRequestsService).DeleteReviewRequests) return pullRequestReviewerFn(ctx, args, "delete", (*gitea_sdk.PullRequestsService).DeleteReviewRequests)
} }
func listPullRequestReviewsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listPullRequestReviewsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -528,8 +521,7 @@ func listPullRequestReviewsFn(ctx context.Context, req mcp.CallToolRequest) (*mc
return to.TextResult(slimReviews(reviews)) return to.TextResult(slimReviews(reviews))
} }
func getPullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getPullRequestReviewFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -560,8 +552,7 @@ func getPullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.
return to.TextResult(slimReview(review)) return to.TextResult(slimReview(review))
} }
func listPullRequestReviewCommentsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listPullRequestReviewCommentsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -612,8 +603,7 @@ func listPullRequestReviewCommentsFn(ctx context.Context, req mcp.CallToolReques
return to.TextResult(slimReviewComments(comments)) return to.TextResult(slimReviewComments(comments))
} }
func createPullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createPullRequestReviewFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -676,8 +666,7 @@ func createPullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(slimReview(review)) return to.TextResult(slimReview(review))
} }
func submitPullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func submitPullRequestReviewFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -719,8 +708,7 @@ func submitPullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(slimReview(review)) return to.TextResult(slimReview(review))
} }
func deletePullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deletePullRequestReviewFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -758,8 +746,7 @@ func deletePullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(successMsg) return to.TextResult(successMsg)
} }
func dismissPullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func dismissPullRequestReviewFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -802,8 +789,7 @@ func dismissPullRequestReviewFn(ctx context.Context, req mcp.CallToolRequest) (*
return to.TextResult(successMsg) return to.TextResult(successMsg)
} }
func replyPullRequestReviewCommentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func replyPullRequestReviewCommentFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -840,16 +826,15 @@ func replyPullRequestReviewCommentFn(ctx context.Context, req mcp.CallToolReques
return to.TextResult(slimReviewComment(comment)) return to.TextResult(slimReviewComment(comment))
} }
func resolveReviewThreadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func resolveReviewThreadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
return setReviewThreadResolvedFn(ctx, req, true) return setReviewThreadResolvedFn(ctx, args, true)
} }
func unresolveReviewThreadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func unresolveReviewThreadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
return setReviewThreadResolvedFn(ctx, req, false) return setReviewThreadResolvedFn(ctx, args, false)
} }
func setReviewThreadResolvedFn(ctx context.Context, req mcp.CallToolRequest, resolved bool) (*mcp.CallToolResult, error) { func setReviewThreadResolvedFn(ctx context.Context, args map[string]any, resolved bool) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -887,8 +872,7 @@ func setReviewThreadResolvedFn(ctx context.Context, req mcp.CallToolRequest, res
return to.TextResult(successMsg) return to.TextResult(successMsg)
} }
func mergePullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func mergePullRequestFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -951,8 +935,7 @@ func mergePullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Call
return to.TextResult(successMsg) return to.TextResult(successMsg)
} }
func editPullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func editPullRequestFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -1026,8 +1009,7 @@ func editPullRequestFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
return to.TextResult(slimPullRequest(pr)) return to.TextResult(slimPullRequest(pr))
} }
func updatePullRequestBranchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func updatePullRequestBranchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -1048,8 +1030,7 @@ func updatePullRequestBranchFn(ctx context.Context, req mcp.CallToolRequest) (*m
return to.TextResult(map[string]any{"message": "branch updated from base"}) return to.TextResult(map[string]any{"message": "branch updated from base"})
} }
func getPullRequestFilesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getPullRequestFilesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -1076,8 +1057,7 @@ func getPullRequestFilesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.C
return to.TextResult(files) return to.TextResult(files)
} }
func getPullRequestStatusFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getPullRequestStatusFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+85 -133
View File
@@ -12,7 +12,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
func Test_editPullRequestFn(t *testing.T) { func Test_editPullRequestFn(t *testing.T) {
@@ -77,19 +77,15 @@ func Test_editPullRequestFn(t *testing.T) {
flag.Version = origVersion flag.Version = origVersion
}() }()
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "owner": owner,
Arguments: map[string]any{ "repo": repo,
"owner": owner, "pull_number": ii.val,
"repo": repo, "title": "WIP: my feature",
"pull_number": ii.val, "state": "open",
"title": "WIP: my feature",
"state": "open",
},
},
} }
result, err := editPullRequestFn(context.Background(), req) result, err := editPullRequestFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("editPullRequestFn() error = %v", err) t.Fatalf("editPullRequestFn() error = %v", err)
} }
@@ -113,7 +109,7 @@ func Test_editPullRequestFn(t *testing.T) {
if len(result.Content) == 0 { if len(result.Content) == 0 {
t.Fatalf("expected content in result") t.Fatalf("expected content in result")
} }
textContent, ok := mcp.AsTextContent(result.Content[0]) textContent, ok := result.Content[0].(*mcp.TextContent)
if !ok { if !ok {
t.Fatalf("expected text content, got %T", result.Content[0]) t.Fatalf("expected text content, got %T", result.Content[0])
} }
@@ -193,21 +189,17 @@ func Test_mergePullRequestFn(t *testing.T) {
flag.Version = origVersion flag.Version = origVersion
}() }()
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "owner": owner,
Arguments: map[string]any{ "repo": repo,
"owner": owner, "pull_number": ii.val,
"repo": repo, "merge_style": "squash",
"pull_number": ii.val, "title": "feat: my squashed commit",
"merge_style": "squash", "message": "Squash merge of PR #5",
"title": "feat: my squashed commit", "delete_branch": true,
"message": "Squash merge of PR #5",
"delete_branch": true,
},
},
} }
result, err := mergePullRequestFn(context.Background(), req) result, err := mergePullRequestFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("mergePullRequestFn() error = %v", err) t.Fatalf("mergePullRequestFn() error = %v", err)
} }
@@ -237,7 +229,7 @@ func Test_mergePullRequestFn(t *testing.T) {
if len(result.Content) == 0 { if len(result.Content) == 0 {
t.Fatalf("expected content in result") t.Fatalf("expected content in result")
} }
textContent, ok := mcp.AsTextContent(result.Content[0]) textContent, ok := result.Content[0].(*mcp.TextContent)
if !ok { if !ok {
t.Fatalf("expected text content, got %T", result.Content[0]) t.Fatalf("expected text content, got %T", result.Content[0])
} }
@@ -306,21 +298,17 @@ func Test_mergePullRequestFn_newParams(t *testing.T) {
flag.Version = origVersion flag.Version = origVersion
}() }()
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "owner": owner,
Arguments: map[string]any{ "repo": repo,
"owner": owner, "pull_number": float64(index),
"repo": repo, "merge_style": "merge",
"pull_number": float64(index), "force_merge": true,
"merge_style": "merge", "merge_when_checks_succeed": true,
"force_merge": true, "head_commit_id": "abc123",
"merge_when_checks_succeed": true,
"head_commit_id": "abc123",
},
},
} }
_, err := mergePullRequestFn(context.Background(), req) _, err := mergePullRequestFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("mergePullRequestFn() error = %v", err) t.Fatalf("mergePullRequestFn() error = %v", err)
} }
@@ -386,22 +374,18 @@ func Test_createPullRequestFn_labels(t *testing.T) {
flag.Version = origVersion flag.Version = origVersion
}() }()
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "owner": owner,
Arguments: map[string]any{ "repo": repo,
"owner": owner, "title": "test",
"repo": repo, "body": "body",
"title": "test", "head": "feature",
"body": "body", "base": "main",
"head": "feature", "labels": []any{float64(1), float64(2)},
"base": "main", "deadline": "2026-06-01T00:00:00Z",
"labels": []any{float64(1), float64(2)},
"deadline": "2026-06-01T00:00:00Z",
},
},
} }
_, err := createPullRequestFn(context.Background(), req) _, err := createPullRequestFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("createPullRequestFn() error = %v", err) t.Fatalf("createPullRequestFn() error = %v", err)
} }
@@ -525,13 +509,7 @@ func Test_createPullRequestFn_draft(t *testing.T) {
args["draft"] = tc.draft args["draft"] = tc.draft
} }
req := mcp.CallToolRequest{ _, err := createPullRequestFn(context.Background(), args)
Params: mcp.CallToolParams{
Arguments: args,
},
}
_, err := createPullRequestFn(context.Background(), req)
if err != nil { if err != nil {
t.Fatalf("createPullRequestFn() error = %v", err) t.Fatalf("createPullRequestFn() error = %v", err)
} }
@@ -630,13 +608,7 @@ func Test_editPullRequestFn_draft(t *testing.T) {
args["draft"] = tc.draft args["draft"] = tc.draft
} }
req := mcp.CallToolRequest{ _, err := editPullRequestFn(context.Background(), args)
Params: mcp.CallToolParams{
Arguments: args,
},
}
_, err := editPullRequestFn(context.Background(), req)
if err != nil { if err != nil {
t.Fatalf("editPullRequestFn() error = %v", err) t.Fatalf("editPullRequestFn() error = %v", err)
} }
@@ -720,18 +692,14 @@ func Test_getPullRequestDiffFn(t *testing.T) {
flag.Version = origVersion flag.Version = origVersion
}() }()
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "owner": owner,
Arguments: map[string]any{ "repo": repo,
"owner": owner, "pull_number": ii.val,
"repo": repo, "binary": true,
"pull_number": ii.val,
"binary": true,
},
},
} }
result, err := getPullRequestDiffFn(context.Background(), req) result, err := getPullRequestDiffFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("getPullRequestDiffFn() error = %v", err) t.Fatalf("getPullRequestDiffFn() error = %v", err)
} }
@@ -758,7 +726,7 @@ func Test_getPullRequestDiffFn(t *testing.T) {
t.Fatalf("expected content in result") t.Fatalf("expected content in result")
} }
textContent, ok := mcp.AsTextContent(result.Content[0]) textContent, ok := result.Content[0].(*mcp.TextContent)
if !ok { if !ok {
t.Fatalf("expected text content, got %T", result.Content[0]) t.Fatalf("expected text content, got %T", result.Content[0])
} }
@@ -807,17 +775,17 @@ func Test_getPullRequestByIndexFn_includesAttachments(t *testing.T) {
flag.Host, flag.Token, flag.Version = server.URL, "", "test" flag.Host, flag.Token, flag.Version = server.URL, "", "test"
defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }() defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }()
req := mcp.CallToolRequest{Params: mcp.CallToolParams{Arguments: map[string]any{ args := map[string]any{
"owner": owner, "repo": repo, "pull_number": float64(index), "owner": owner, "repo": repo, "pull_number": float64(index),
}}} }
res, err := getPullRequestByIndexFn(context.Background(), req) res, err := getPullRequestByIndexFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("getPullRequestByIndexFn() error = %v", err) t.Fatalf("getPullRequestByIndexFn() error = %v", err)
} }
if res.IsError { if res.IsError {
t.Fatalf("unexpected error result: %v", res.Content) t.Fatalf("unexpected error result: %v", res.Content)
} }
body := res.Content[0].(mcp.TextContent).Text body := res.Content[0].(*mcp.TextContent).Text
if !strings.Contains(body, `[shot.png](https://example/shot.png)`) { if !strings.Contains(body, `[shot.png](https://example/shot.png)`) {
t.Fatalf("expected attachment markdown inlined in body, got: %s", body) t.Fatalf("expected attachment markdown inlined in body, got: %s", body)
} }
@@ -855,14 +823,14 @@ func Test_getPullRequestByIndexFn_emptyAssetsLeavesBody(t *testing.T) {
flag.Host, flag.Token, flag.Version = server.URL, "", "test" flag.Host, flag.Token, flag.Version = server.URL, "", "test"
defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }() defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }()
req := mcp.CallToolRequest{Params: mcp.CallToolParams{Arguments: map[string]any{ args := map[string]any{
"owner": owner, "repo": repo, "pull_number": float64(index), "owner": owner, "repo": repo, "pull_number": float64(index),
}}} }
res, err := getPullRequestByIndexFn(context.Background(), req) res, err := getPullRequestByIndexFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("getPullRequestByIndexFn() error = %v", err) t.Fatalf("getPullRequestByIndexFn() error = %v", err)
} }
body := res.Content[0].(mcp.TextContent).Text body := res.Content[0].(*mcp.TextContent).Text
if !strings.Contains(body, `"body":"plain body"`) { if !strings.Contains(body, `"body":"plain body"`) {
t.Fatalf("expected body unchanged when assets are empty, got: %s", body) t.Fatalf("expected body unchanged when assets are empty, got: %s", body)
} }
@@ -899,17 +867,17 @@ func Test_getPullRequestByIndexFn_assetsFailureNonFatal(t *testing.T) {
flag.Host, flag.Token, flag.Version = server.URL, "", "test" flag.Host, flag.Token, flag.Version = server.URL, "", "test"
defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }() defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }()
req := mcp.CallToolRequest{Params: mcp.CallToolParams{Arguments: map[string]any{ args := map[string]any{
"owner": owner, "repo": repo, "pull_number": float64(index), "owner": owner, "repo": repo, "pull_number": float64(index),
}}} }
res, err := getPullRequestByIndexFn(context.Background(), req) res, err := getPullRequestByIndexFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("getPullRequestByIndexFn() error = %v", err) t.Fatalf("getPullRequestByIndexFn() error = %v", err)
} }
if res.IsError { if res.IsError {
t.Fatalf("assets fetch failure should not fail the PR fetch: %v", res.Content) t.Fatalf("assets fetch failure should not fail the PR fetch: %v", res.Content)
} }
body := res.Content[0].(mcp.TextContent).Text body := res.Content[0].(*mcp.TextContent).Text
if !strings.Contains(body, `"plain body"`) { if !strings.Contains(body, `"plain body"`) {
t.Fatalf("expected PR body preserved when assets fail, got: %s", body) t.Fatalf("expected PR body preserved when assets fail, got: %s", body)
} }
@@ -954,18 +922,14 @@ func Test_closePullRequestFn(t *testing.T) {
flag.Token = "test-token" flag.Token = "test-token"
t.Cleanup(func() { flag.Host = origHost; flag.Token = origToken }) t.Cleanup(func() { flag.Host = origHost; flag.Token = origToken })
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "method": "close",
Arguments: map[string]any{ "owner": owner,
"method": "close", "repo": repo,
"owner": owner, "pull_number": float64(index),
"repo": repo,
"pull_number": float64(index),
},
},
} }
result, err := closePullRequestFn(context.Background(), req) result, err := closePullRequestFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("closePullRequestFn() error = %v", err) t.Fatalf("closePullRequestFn() error = %v", err)
} }
@@ -1018,18 +982,14 @@ func Test_reopenPullRequestFn(t *testing.T) {
flag.Token = "test-token" flag.Token = "test-token"
t.Cleanup(func() { flag.Host = origHost; flag.Token = origToken }) t.Cleanup(func() { flag.Host = origHost; flag.Token = origToken })
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "method": "reopen",
Arguments: map[string]any{ "owner": owner,
"method": "reopen", "repo": repo,
"owner": owner, "pull_number": float64(index),
"repo": repo,
"pull_number": float64(index),
},
},
} }
result, err := reopenPullRequestFn(context.Background(), req) result, err := reopenPullRequestFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("reopenPullRequestFn() error = %v", err) t.Fatalf("reopenPullRequestFn() error = %v", err)
} }
@@ -1103,20 +1063,16 @@ func Test_pullRequestReviewWriteFn_comments(t *testing.T) {
_, _ = w.Write([]byte(`{"id":43,"body":"sure","path":"main.go","position":3}`)) _, _ = w.Write([]byte(`{"id":43,"body":"sure","path":"main.go","position":3}`))
}) })
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "method": tc.method,
Arguments: map[string]any{ "owner": owner,
"method": tc.method, "repo": repo,
"owner": owner, "pull_number": float64(index),
"repo": repo, "comment_id": float64(commentID),
"pull_number": float64(index), "body": "sure",
"comment_id": float64(commentID),
"body": "sure",
},
},
} }
result, err := pullRequestReviewWriteFn(context.Background(), req) result, err := pullRequestReviewWriteFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("pullRequestReviewWriteFn() error = %v", err) t.Fatalf("pullRequestReviewWriteFn() error = %v", err)
} }
@@ -1162,18 +1118,14 @@ func Test_listPullRequestReviewCommentsFn_allReviews(t *testing.T) {
} }
}) })
req := mcp.CallToolRequest{ args := map[string]any{
Params: mcp.CallToolParams{ "method": "get_review_comments",
Arguments: map[string]any{ "owner": owner,
"method": "get_review_comments", "repo": repo,
"owner": owner, "pull_number": float64(index),
"repo": repo,
"pull_number": float64(index),
},
},
} }
result, err := pullRequestReadFn(context.Background(), req) result, err := pullRequestReadFn(context.Background(), args)
if err != nil { if err != nil {
t.Fatalf("pullRequestReadFn() error = %v", err) t.Fatalf("pullRequestReadFn() error = %v", err)
} }
+27 -31
View File
@@ -11,8 +11,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
// BranchTool holds the branch-related tools (scope "branch"). // BranchTool holds the branch-related tools (scope "branch").
@@ -25,53 +24,52 @@ const (
) )
var ( var (
CreateBranchTool = mcp.NewTool( CreateBranchTool = tool.NewDefinition(
CreateBranchToolName, CreateBranchToolName,
mcp.WithDescription("Create a new branch in a repository, optionally from a specific source branch (defaults to the repository's default branch)."), "Create a new branch in a repository, optionally from a specific source branch (defaults to the repository's default branch).",
mcp.WithToolAnnotation(annotation.Write("Create a new branch")), annotation.Write("Create a new branch"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("branch", mcp.Required()), tool.String("branch", tool.Required()),
mcp.WithString("old_branch", mcp.Description("source branch (default: repo default)")), tool.String("old_branch", tool.Description("source branch (default: repo default)")),
) )
DeleteBranchTool = mcp.NewTool( DeleteBranchTool = tool.NewDefinition(
DeleteBranchToolName, DeleteBranchToolName,
mcp.WithDescription("Permanently delete a branch from a repository. This action is destructive and cannot be undone."), "Permanently delete a branch from a repository. This action is destructive and cannot be undone.",
mcp.WithToolAnnotation(annotation.Destructive("Delete a branch")), annotation.Destructive("Delete a branch"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("branch", mcp.Required()), tool.String("branch", tool.Required()),
) )
ListBranchesTool = mcp.NewTool( ListBranchesTool = tool.NewDefinition(
ListBranchesToolName, ListBranchesToolName,
mcp.WithDescription("List all branches in a repository, paginated."), "List all branches in a repository, paginated.",
mcp.WithToolAnnotation(annotation.ReadOnly("List repository branches")), annotation.ReadOnly("List repository branches"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
) )
func init() { func init() {
BranchTool.RegisterWrite(server.ServerTool{ BranchTool.RegisterWrite(tool.ServerTool{
Tool: CreateBranchTool, Tool: CreateBranchTool,
Handler: CreateBranchFn, Handler: CreateBranchFn,
}) })
BranchTool.RegisterWrite(server.ServerTool{ BranchTool.RegisterWrite(tool.ServerTool{
Tool: DeleteBranchTool, Tool: DeleteBranchTool,
Handler: DeleteBranchFn, Handler: DeleteBranchFn,
}) })
BranchTool.RegisterRead(server.ServerTool{ BranchTool.RegisterRead(tool.ServerTool{
Tool: ListBranchesTool, Tool: ListBranchesTool,
Handler: ListBranchesFn, Handler: ListBranchesFn,
}) })
} }
func CreateBranchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func CreateBranchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -101,8 +99,7 @@ func CreateBranchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTool
return to.TextResult("Branch Created") return to.TextResult("Branch Created")
} }
func DeleteBranchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func DeleteBranchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -127,8 +124,7 @@ func DeleteBranchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTool
return to.TextResult("Branch Deleted") return to.TextResult("Branch Deleted")
} }
func ListBranchesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func ListBranchesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+20 -23
View File
@@ -11,8 +11,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
// CommitTool holds the commit-related tools (scope "commit"). // CommitTool holds the commit-related tools (scope "commit").
@@ -24,41 +23,40 @@ const (
) )
var ( var (
ListRepoCommitsTool = mcp.NewTool( ListRepoCommitsTool = tool.NewDefinition(
ListRepoCommitsToolName, ListRepoCommitsToolName,
mcp.WithDescription("List commits in a repository, optionally starting from a specific branch or SHA and filtered to commits touching a given file path."), "List commits in a repository, optionally starting from a specific branch or SHA and filtered to commits touching a given file path.",
mcp.WithToolAnnotation(annotation.ReadOnly("List repository commits")), annotation.ReadOnly("List repository commits"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("sha", mcp.Description("starting SHA or branch")), tool.String("sha", tool.Description("starting SHA or branch")),
mcp.WithString("path", mcp.Description("only commits touching this path")), tool.String("path", tool.Description("only commits touching this path")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1), mcp.Min(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1), tool.Minimum(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30), mcp.Min(1)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30), tool.Minimum(1)),
) )
GetCommitTool = mcp.NewTool( GetCommitTool = tool.NewDefinition(
GetCommitToolName, GetCommitToolName,
mcp.WithDescription("Get details for a single commit in a repository by its SHA."), "Get details for a single commit in a repository by its SHA.",
mcp.WithToolAnnotation(annotation.ReadOnly("Get commit details")), annotation.ReadOnly("Get commit details"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("sha", mcp.Required()), tool.String("sha", tool.Required()),
) )
) )
func init() { func init() {
CommitTool.RegisterRead(server.ServerTool{ CommitTool.RegisterRead(tool.ServerTool{
Tool: ListRepoCommitsTool, Tool: ListRepoCommitsTool,
Handler: ListRepoCommitsFn, Handler: ListRepoCommitsFn,
}) })
CommitTool.RegisterRead(server.ServerTool{ CommitTool.RegisterRead(tool.ServerTool{
Tool: GetCommitTool, Tool: GetCommitTool,
Handler: GetCommitFn, Handler: GetCommitFn,
}) })
} }
func ListRepoCommitsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func ListRepoCommitsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -89,8 +87,7 @@ func ListRepoCommitsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
return to.TextResult(slimCommits(commits)) return to.TextResult(slimCommits(commits))
} }
func GetCommitFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetCommitFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+44 -49
View File
@@ -15,8 +15,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
// FileTool holds the file-related tools (scope "file"). // FileTool holds the file-related tools (scope "file").
@@ -30,68 +29,68 @@ const (
) )
var ( var (
GetFileContentTool = mcp.NewTool( GetFileContentTool = tool.NewDefinition(
GetFileToolName, GetFileToolName,
mcp.WithDescription("Get file content and metadata"), "Get file content and metadata",
mcp.WithToolAnnotation(annotation.ReadOnly("Get file content")), annotation.ReadOnly("Get file content"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("ref", mcp.Required(), mcp.Description("branch, tag, or commit SHA")), tool.String("ref", tool.Required(), tool.Description("branch, tag, or commit SHA")),
mcp.WithString("path", mcp.Required()), tool.String("path", tool.Required()),
mcp.WithBoolean("withLines", mcp.Description("return numbered lines")), tool.Boolean("withLines", tool.Description("return numbered lines")),
) )
GetDirContentTool = mcp.NewTool( GetDirContentTool = tool.NewDefinition(
GetDirToolName, GetDirToolName,
mcp.WithDescription("List the entries (files and subdirectories) in a repository directory at a given ref (branch, tag, or commit SHA)."), "List the entries (files and subdirectories) in a repository directory at a given ref (branch, tag, or commit SHA).",
mcp.WithToolAnnotation(annotation.ReadOnly("Get directory contents")), annotation.ReadOnly("Get directory contents"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("ref", mcp.Required(), mcp.Description("branch, tag, or commit SHA")), tool.String("ref", tool.Required(), tool.Description("branch, tag, or commit SHA")),
mcp.WithString("path", mcp.Required()), tool.String("path", tool.Required()),
) )
CreateOrUpdateFileTool = mcp.NewTool( CreateOrUpdateFileTool = tool.NewDefinition(
CreateOrUpdateFileToolName, CreateOrUpdateFileToolName,
mcp.WithDescription("Create or update a file (provide sha to update an existing file)."), "Create or update a file (provide sha to update an existing file).",
mcp.WithToolAnnotation(annotation.Write("Create or update a file")), annotation.Write("Create or update a file"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("path", mcp.Required()), tool.String("path", tool.Required()),
mcp.WithString("content", mcp.Required()), tool.String("content", tool.Required()),
mcp.WithString("message", mcp.Required(), mcp.Description("commit message")), tool.String("message", tool.Required(), tool.Description("commit message")),
mcp.WithString("branch_name", mcp.Required()), tool.String("branch_name", tool.Required()),
mcp.WithString("sha", mcp.Description("existing file SHA (omit to create)")), tool.String("sha", tool.Description("existing file SHA (omit to create)")),
mcp.WithString("new_branch_name", mcp.Description("new branch (create only)")), tool.String("new_branch_name", tool.Description("new branch (create only)")),
) )
DeleteFileTool = mcp.NewTool( DeleteFileTool = tool.NewDefinition(
DeleteFileToolName, DeleteFileToolName,
mcp.WithDescription("Delete a file from a repository by committing the removal to a branch. Requires the file's current SHA and a commit message."), "Delete a file from a repository by committing the removal to a branch. Requires the file's current SHA and a commit message.",
mcp.WithToolAnnotation(annotation.Destructive("Delete a file")), annotation.Destructive("Delete a file"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("path", mcp.Required()), tool.String("path", tool.Required()),
mcp.WithString("message", mcp.Required(), mcp.Description("commit message")), tool.String("message", tool.Required(), tool.Description("commit message")),
mcp.WithString("branch_name", mcp.Required()), tool.String("branch_name", tool.Required()),
mcp.WithString("sha", mcp.Required()), tool.String("sha", tool.Required()),
) )
) )
func init() { func init() {
FileTool.RegisterRead(server.ServerTool{ FileTool.RegisterRead(tool.ServerTool{
Tool: GetFileContentTool, Tool: GetFileContentTool,
Handler: GetFileContentFn, Handler: GetFileContentFn,
}) })
FileTool.RegisterRead(server.ServerTool{ FileTool.RegisterRead(tool.ServerTool{
Tool: GetDirContentTool, Tool: GetDirContentTool,
Handler: GetDirContentFn, Handler: GetDirContentFn,
}) })
FileTool.RegisterWrite(server.ServerTool{ FileTool.RegisterWrite(tool.ServerTool{
Tool: CreateOrUpdateFileTool, Tool: CreateOrUpdateFileTool,
Handler: CreateOrUpdateFileFn, Handler: CreateOrUpdateFileFn,
}) })
FileTool.RegisterWrite(server.ServerTool{ FileTool.RegisterWrite(tool.ServerTool{
Tool: DeleteFileTool, Tool: DeleteFileTool,
Handler: DeleteFileFn, Handler: DeleteFileFn,
}) })
@@ -102,8 +101,7 @@ type ContentLine struct {
Content string `json:"content"` Content string `json:"content"`
} }
func GetFileContentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetFileContentFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -165,8 +163,7 @@ func GetFileContentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(slimContents(content)) return to.TextResult(slimContents(content))
} }
func GetDirContentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetDirContentFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -191,8 +188,7 @@ func GetDirContentFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToo
return to.TextResult(slimDirEntries(content)) return to.TextResult(slimDirEntries(content))
} }
func CreateOrUpdateFileFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func CreateOrUpdateFileFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -250,8 +246,7 @@ func CreateOrUpdateFileFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Ca
return to.TextResult("Create file success") return to.TextResult("Create file success")
} }
func DeleteFileFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func DeleteFileFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+48 -54
View File
@@ -11,8 +11,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
// ReleaseTool holds the release-related tools (scope "release"). // ReleaseTool holds the release-related tools (scope "release").
@@ -27,84 +26,83 @@ const (
) )
var ( var (
CreateReleaseTool = mcp.NewTool( CreateReleaseTool = tool.NewDefinition(
CreateReleaseToolName, CreateReleaseToolName,
mcp.WithDescription("Create a new release in a repository from a tag, optionally marking it as a draft or pre-release."), "Create a new release in a repository from a tag, optionally marking it as a draft or pre-release.",
mcp.WithToolAnnotation(annotation.Write("Create a release")), annotation.Write("Create a release"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("tag_name", mcp.Required()), tool.String("tag_name", tool.Required()),
mcp.WithString("target", mcp.Required(), mcp.Description("commitish")), tool.String("target", tool.Required(), tool.Description("commitish")),
mcp.WithString("title", mcp.Required()), tool.String("title", tool.Required()),
mcp.WithBoolean("is_draft"), tool.Boolean("is_draft"),
mcp.WithBoolean("is_pre_release"), tool.Boolean("is_pre_release"),
mcp.WithString("body"), tool.String("body"),
) )
DeleteReleaseTool = mcp.NewTool( DeleteReleaseTool = tool.NewDefinition(
DeleteReleaseToolName, DeleteReleaseToolName,
mcp.WithDescription("Delete a release from a repository by its numeric ID. This action is destructive and cannot be undone."), "Delete a release from a repository by its numeric ID. This action is destructive and cannot be undone.",
mcp.WithToolAnnotation(annotation.Destructive("Delete a release")), annotation.Destructive("Delete a release"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("id", mcp.Required()), tool.Number("id", tool.Required()),
) )
GetReleaseTool = mcp.NewTool( GetReleaseTool = tool.NewDefinition(
GetReleaseToolName, GetReleaseToolName,
mcp.WithDescription("Get a release by ID"), "Get a release by ID",
mcp.WithToolAnnotation(annotation.ReadOnly("Get release details")), annotation.ReadOnly("Get release details"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("id", mcp.Required()), tool.Number("id", tool.Required()),
) )
GetLatestReleaseTool = mcp.NewTool( GetLatestReleaseTool = tool.NewDefinition(
GetLatestReleaseToolName, GetLatestReleaseToolName,
mcp.WithDescription("Get the most recent published (non-draft) release in a repository."), "Get the most recent published (non-draft) release in a repository.",
mcp.WithToolAnnotation(annotation.ReadOnly("Get latest release")), annotation.ReadOnly("Get latest release"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
) )
ListReleasesTool = mcp.NewTool( ListReleasesTool = tool.NewDefinition(
ListReleasesToolName, ListReleasesToolName,
mcp.WithDescription("List releases in a repository, optionally filtered to drafts or pre-releases."), "List releases in a repository, optionally filtered to drafts or pre-releases.",
mcp.WithToolAnnotation(annotation.ReadOnly("List releases")), annotation.ReadOnly("List releases"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithBoolean("is_draft"), tool.Boolean("is_draft"),
mcp.WithBoolean("is_pre_release"), tool.Boolean("is_pre_release"),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1), mcp.Min(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1), tool.Minimum(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(20), mcp.Min(1)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(20), tool.Minimum(1)),
) )
) )
func init() { func init() {
ReleaseTool.RegisterWrite(server.ServerTool{ ReleaseTool.RegisterWrite(tool.ServerTool{
Tool: CreateReleaseTool, Tool: CreateReleaseTool,
Handler: CreateReleaseFn, Handler: CreateReleaseFn,
}) })
ReleaseTool.RegisterWrite(server.ServerTool{ ReleaseTool.RegisterWrite(tool.ServerTool{
Tool: DeleteReleaseTool, Tool: DeleteReleaseTool,
Handler: DeleteReleaseFn, Handler: DeleteReleaseFn,
}) })
ReleaseTool.RegisterRead(server.ServerTool{ ReleaseTool.RegisterRead(tool.ServerTool{
Tool: GetReleaseTool, Tool: GetReleaseTool,
Handler: GetReleaseFn, Handler: GetReleaseFn,
}) })
ReleaseTool.RegisterRead(server.ServerTool{ ReleaseTool.RegisterRead(tool.ServerTool{
Tool: GetLatestReleaseTool, Tool: GetLatestReleaseTool,
Handler: GetLatestReleaseFn, Handler: GetLatestReleaseFn,
}) })
ReleaseTool.RegisterRead(server.ServerTool{ ReleaseTool.RegisterRead(tool.ServerTool{
Tool: ListReleasesTool, Tool: ListReleasesTool,
Handler: ListReleasesFn, Handler: ListReleasesFn,
}) })
} }
func CreateReleaseFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func CreateReleaseFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -148,8 +146,7 @@ func CreateReleaseFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToo
return to.TextResult("Release Created") return to.TextResult("Release Created")
} }
func DeleteReleaseFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func DeleteReleaseFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -175,8 +172,7 @@ func DeleteReleaseFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToo
return to.TextResult("Release deleted successfully") return to.TextResult("Release deleted successfully")
} }
func GetReleaseFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetReleaseFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -202,8 +198,7 @@ func GetReleaseFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolRe
return to.TextResult(slimRelease(release)) return to.TextResult(slimRelease(release))
} }
func GetLatestReleaseFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetLatestReleaseFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -225,8 +220,7 @@ func GetLatestReleaseFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Call
return to.TextResult(slimRelease(release)) return to.TextResult(slimRelease(release))
} }
func ListReleasesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func ListReleasesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+46 -49
View File
@@ -12,8 +12,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("repository") var Tool = tool.New("repository")
@@ -26,74 +25,73 @@ const (
) )
var ( var (
CreateRepoTool = mcp.NewTool( CreateRepoTool = tool.NewDefinition(
CreateRepoToolName, CreateRepoToolName,
mcp.WithDescription("Create a new Git repository, optionally under an organization (defaults to the authenticated user's account), with options for visibility, template, license, .gitignore, and initial README."), "Create a new Git repository, optionally under an organization (defaults to the authenticated user's account), with options for visibility, template, license, .gitignore, and initial README.",
mcp.WithToolAnnotation(annotation.Write("Create a new repository")), annotation.Write("Create a new repository"),
mcp.WithString("name", mcp.Required()), tool.String("name", tool.Required()),
mcp.WithString("description"), tool.String("description"),
mcp.WithBoolean("private"), tool.Boolean("private"),
mcp.WithString("issue_labels"), tool.String("issue_labels"),
mcp.WithBoolean("auto_init"), tool.Boolean("auto_init"),
mcp.WithBoolean("template"), tool.Boolean("template"),
mcp.WithString("gitignores"), tool.String("gitignores"),
mcp.WithString("license"), tool.String("license"),
mcp.WithString("readme"), tool.String("readme"),
mcp.WithString("default_branch"), tool.String("default_branch"),
mcp.WithString("trust_model", mcp.Enum("default", "collaborator", "committer", "collaboratorcommitter")), tool.String("trust_model", tool.Enum("default", "collaborator", "committer", "collaboratorcommitter")),
mcp.WithString("object_format_name", mcp.Enum("sha1", "sha256")), tool.String("object_format_name", tool.Enum("sha1", "sha256")),
mcp.WithString("organization", mcp.Description("defaults to personal account")), tool.String("organization", tool.Description("defaults to personal account")),
) )
ForkRepoTool = mcp.NewTool( ForkRepoTool = tool.NewDefinition(
ForkRepoToolName, ForkRepoToolName,
mcp.WithDescription("Fork an existing repository into the authenticated user's account or a target organization, optionally under a new name."), "Fork an existing repository into the authenticated user's account or a target organization, optionally under a new name.",
mcp.WithToolAnnotation(annotation.Write("Fork a repository")), annotation.Write("Fork a repository"),
mcp.WithString("user", mcp.Required(), mcp.Description("owner of source repo")), tool.String("user", tool.Required(), tool.Description("owner of source repo")),
mcp.WithString("repo", mcp.Required()), tool.String("repo", tool.Required()),
mcp.WithString("organization", mcp.Description("target org")), tool.String("organization", tool.Description("target org")),
mcp.WithString("name", mcp.Description("fork name")), tool.String("name", tool.Description("fork name")),
) )
ListMyReposTool = mcp.NewTool( ListMyReposTool = tool.NewDefinition(
ListMyReposToolName, ListMyReposToolName,
mcp.WithDescription("List repositories owned by the authenticated user."), "List repositories owned by the authenticated user.",
mcp.WithToolAnnotation(annotation.ReadOnly("List my repositories")), annotation.ReadOnly("List my repositories"),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1), mcp.Min(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1), tool.Minimum(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30), mcp.Min(1)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30), tool.Minimum(1)),
) )
ListOrgReposTool = mcp.NewTool( ListOrgReposTool = tool.NewDefinition(
ListOrgReposToolName, ListOrgReposToolName,
mcp.WithDescription("List repositories belonging to an organization."), "List repositories belonging to an organization.",
mcp.WithToolAnnotation(annotation.ReadOnly("List organization repositories")), annotation.ReadOnly("List organization repositories"),
mcp.WithString("org", mcp.Required()), tool.String("org", tool.Required()),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1), mcp.Min(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1), tool.Minimum(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(100), mcp.Min(1)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(100), tool.Minimum(1)),
) )
) )
func init() { func init() {
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: CreateRepoTool, Tool: CreateRepoTool,
Handler: CreateRepoFn, Handler: CreateRepoFn,
}) })
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: ForkRepoTool, Tool: ForkRepoTool,
Handler: ForkRepoFn, Handler: ForkRepoFn,
}) })
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: ListMyReposTool, Tool: ListMyReposTool,
Handler: ListMyReposFn, Handler: ListMyReposFn,
}) })
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: ListOrgReposTool, Tool: ListOrgReposTool,
Handler: ListOrgReposFn, Handler: ListOrgReposFn,
}) })
} }
func CreateRepoFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func CreateRepoFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
name, err := params.GetString(args, "name") name, err := params.GetString(args, "name")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -145,8 +143,7 @@ func CreateRepoFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolRe
return to.TextResult(slim.Repo(repo)) return to.TextResult(slim.Repo(repo))
} }
func ForkRepoFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func ForkRepoFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
user, err := params.GetString(args, "user") user, err := params.GetString(args, "user")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -170,8 +167,8 @@ func ForkRepoFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResu
return to.TextResult("Fork success") return to.TextResult("Fork success")
} }
func ListMyReposFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func ListMyReposFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
opt := gitea_sdk.ListReposOptions{ opt := gitea_sdk.ListReposOptions{
ListOptions: gitea_sdk.ListOptions{ ListOptions: gitea_sdk.ListOptions{
Page: page, Page: page,
@@ -190,12 +187,12 @@ func ListMyReposFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolR
return to.TextResult(slim.Repos(repos)) return to.TextResult(slim.Repos(repos))
} }
func ListOrgReposFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func ListOrgReposFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 100) page, pageSize := params.GetPagination(args, 100)
opt := gitea_sdk.ListOrgReposOptions{ opt := gitea_sdk.ListOrgReposOptions{
ListOptions: gitea_sdk.ListOptions{ ListOptions: gitea_sdk.ListOptions{
Page: page, Page: page,
+36 -41
View File
@@ -11,8 +11,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
// TagTool holds the tag-related tools (scope "tag"). // TagTool holds the tag-related tools (scope "tag").
@@ -26,67 +25,66 @@ const (
) )
var ( var (
CreateTagTool = mcp.NewTool( CreateTagTool = tool.NewDefinition(
CreateTagToolName, CreateTagToolName,
mcp.WithDescription("Create a new Git tag in a repository at a target commit, branch, or existing tag, with an optional annotation message."), "Create a new Git tag in a repository at a target commit, branch, or existing tag, with an optional annotation message.",
mcp.WithToolAnnotation(annotation.Write("Create a tag")), annotation.Write("Create a tag"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("tag_name", mcp.Required()), tool.String("tag_name", tool.Required()),
mcp.WithString("target", mcp.Description("commitish")), tool.String("target", tool.Description("commitish")),
mcp.WithString("message", mcp.Description("tag message")), tool.String("message", tool.Description("tag message")),
) )
DeleteTagTool = mcp.NewTool( DeleteTagTool = tool.NewDefinition(
DeleteTagToolName, DeleteTagToolName,
mcp.WithDescription("Permanently delete a tag from a repository. This action is destructive and cannot be undone."), "Permanently delete a tag from a repository. This action is destructive and cannot be undone.",
mcp.WithToolAnnotation(annotation.Destructive("Delete a tag")), annotation.Destructive("Delete a tag"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("tag_name", mcp.Required()), tool.String("tag_name", tool.Required()),
) )
GetTagTool = mcp.NewTool( GetTagTool = tool.NewDefinition(
GetTagToolName, GetTagToolName,
mcp.WithDescription("Get details for a single tag in a repository by name."), "Get details for a single tag in a repository by name.",
mcp.WithToolAnnotation(annotation.ReadOnly("Get tag details")), annotation.ReadOnly("Get tag details"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("tag_name", mcp.Required()), tool.String("tag_name", tool.Required()),
) )
ListTagsTool = mcp.NewTool( ListTagsTool = tool.NewDefinition(
ListTagsToolName, ListTagsToolName,
mcp.WithDescription("List all tags in a repository, paginated."), "List all tags in a repository, paginated.",
mcp.WithToolAnnotation(annotation.ReadOnly("List tags")), annotation.ReadOnly("List tags"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1), mcp.Min(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1), tool.Minimum(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(20), mcp.Min(1)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(20), tool.Minimum(1)),
) )
) )
func init() { func init() {
TagTool.RegisterWrite(server.ServerTool{ TagTool.RegisterWrite(tool.ServerTool{
Tool: CreateTagTool, Tool: CreateTagTool,
Handler: CreateTagFn, Handler: CreateTagFn,
}) })
TagTool.RegisterWrite(server.ServerTool{ TagTool.RegisterWrite(tool.ServerTool{
Tool: DeleteTagTool, Tool: DeleteTagTool,
Handler: DeleteTagFn, Handler: DeleteTagFn,
}) })
TagTool.RegisterRead(server.ServerTool{ TagTool.RegisterRead(tool.ServerTool{
Tool: GetTagTool, Tool: GetTagTool,
Handler: GetTagFn, Handler: GetTagFn,
}) })
TagTool.RegisterRead(server.ServerTool{ TagTool.RegisterRead(tool.ServerTool{
Tool: ListTagsTool, Tool: ListTagsTool,
Handler: ListTagsFn, Handler: ListTagsFn,
}) })
} }
func CreateTagFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func CreateTagFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -118,8 +116,7 @@ func CreateTagFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolRes
return to.TextResult("Tag Created") return to.TextResult("Tag Created")
} }
func DeleteTagFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func DeleteTagFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -145,8 +142,7 @@ func DeleteTagFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolRes
return to.TextResult("Tag deleted") return to.TextResult("Tag deleted")
} }
func GetTagFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetTagFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -172,8 +168,7 @@ func GetTagFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult
return to.TextResult(slimTag(tag)) return to.TextResult(slimTag(tag))
} }
func ListTagsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func ListTagsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+13 -14
View File
@@ -8,37 +8,36 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/gitea" "gitea.com/gitea/gitea-mcp/pkg/gitea"
"gitea.com/gitea/gitea-mcp/pkg/params" "gitea.com/gitea/gitea-mcp/pkg/params"
"gitea.com/gitea/gitea-mcp/pkg/to" "gitea.com/gitea/gitea-mcp/pkg/to"
"gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
const ( const (
GetRepoTreeToolName = "get_repository_tree" GetRepoTreeToolName = "get_repository_tree"
) )
var GetRepoTreeTool = mcp.NewTool( var GetRepoTreeTool = tool.NewDefinition(
GetRepoTreeToolName, GetRepoTreeToolName,
mcp.WithDescription("Get the file tree of a repository at a given ref (SHA, branch, or tag), optionally recursively."), "Get the file tree of a repository at a given ref (SHA, branch, or tag), optionally recursively.",
mcp.WithToolAnnotation(annotation.ReadOnly("Get repository file tree")), annotation.ReadOnly("Get repository file tree"),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("tree_sha", mcp.Required(), mcp.Description("SHA, branch, or tag")), tool.String("tree_sha", tool.Required(), tool.Description("SHA, branch, or tag")),
mcp.WithBoolean("recursive"), tool.Boolean("recursive"),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: GetRepoTreeTool, Tool: GetRepoTreeTool,
Handler: GetRepoTreeFn, Handler: GetRepoTreeFn,
}) })
} }
func GetRepoTreeFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetRepoTreeFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+3 -1
View File
@@ -44,8 +44,10 @@ func TestSlimTreeNil(t *testing.T) {
} }
func TestGetRepoTreeToolRequired(t *testing.T) { func TestGetRepoTreeToolRequired(t *testing.T) {
inputSchema := GetRepoTreeTool.InputSchema.(map[string]any)
required, _ := inputSchema["required"].([]string)
for _, field := range []string{"owner", "repo", "tree_sha"} { for _, field := range []string{"owner", "repo", "tree_sha"} {
if !slices.Contains(GetRepoTreeTool.InputSchema.Required, field) { if !slices.Contains(required, field) {
t.Errorf("expected %q to be required", field) t.Errorf("expected %q to be required", field)
} }
} }
+404
View File
@@ -0,0 +1,404 @@
package operation
import (
"context"
"errors"
"fmt"
"io"
"net/http"
"net/http/httptest"
"os"
"os/exec"
"path/filepath"
"strings"
"sync"
"testing"
"time"
mcpContext "gitea.com/gitea/gitea-mcp/pkg/context"
"gitea.com/gitea/gitea-mcp/pkg/flag"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
// Pin negotiated versions so SDK upgrades require compatibility review.
const (
testServerVersion = "test-version"
expectedProtocolVersion = "2026-07-28"
expectedStatefulHTTPProtocolVersion = "2025-11-25"
)
func exposeAllTools(t *testing.T) {
t.Helper()
originalReadOnly := flag.ReadOnly
originalAllowedTools := flag.AllowedTools
originalAllowedScopes := flag.AllowedScopes
originalVersion := flag.Version
t.Cleanup(func() {
flag.ReadOnly = originalReadOnly
flag.AllowedTools = originalAllowedTools
flag.AllowedScopes = originalAllowedScopes
flag.Version = originalVersion
})
flag.ReadOnly = false
flag.AllowedTools = nil
flag.AllowedScopes = nil
flag.Version = testServerVersion
}
// registeredToolCount is what the registry exposes under the current flags, so
// the transport assertions track tool additions without being edited.
func registeredToolCount() int {
count := 0
for _, domain := range domainTools {
count += len(domain.Tools())
}
return count
}
// stdioCommandEnvironment removes variables that override subprocess flags.
func stdioCommandEnvironment() []string {
environment := os.Environ()
filtered := make([]string, 0, len(environment))
for _, entry := range environment {
name, _, _ := strings.Cut(entry, "=")
switch name {
case "GITEA_READONLY", "GITEA_SCOPES", "GITEA_TOOLS", "MCP_MODE":
continue
}
filtered = append(filtered, entry)
}
return filtered
}
func textContent(t *testing.T, result *mcp.CallToolResult) string {
t.Helper()
if len(result.Content) != 1 {
t.Fatalf("content count = %d, want 1", len(result.Content))
}
content, ok := result.Content[0].(*mcp.TextContent)
if !ok {
t.Fatalf("content type = %T, want *mcp.TextContent", result.Content[0])
}
return content.Text
}
// listAndCallVersion is the round trip every transport must support. wantText
// differs per transport: the stdio subprocess resolves its version from the VCS
// build info (main.go:14), so only the in-process servers have a known one.
func listAndCallVersion(ctx context.Context, t *testing.T, session *mcp.ClientSession, wantText string) {
t.Helper()
result, err := session.ListTools(ctx, nil)
if err != nil {
t.Fatalf("ListTools() error = %v", err)
}
if want := registeredToolCount(); len(result.Tools) != want {
t.Fatalf("ListTools() count = %d, want %d", len(result.Tools), want)
}
callResult, err := session.CallTool(ctx, &mcp.CallToolParams{
Name: "get_gitea_mcp_server_version",
})
if err != nil {
t.Fatalf("CallTool() error = %v", err)
}
if got := textContent(t, callResult); !strings.Contains(got, wantText) {
t.Errorf("version tool result = %q, want it to contain %q", got, wantText)
}
}
func TestOfficialSDKInMemory(t *testing.T) {
exposeAllTools(t)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
serverTransport, clientTransport := mcp.NewInMemoryTransports()
server := newMCPServer(testServerVersion)
RegisterTool(server)
serverDone := make(chan error, 1)
go func() {
serverDone <- server.Run(ctx, serverTransport)
}()
client := mcp.NewClient(&mcp.Implementation{Name: "gitea-mcp-test", Version: "1"}, nil)
session, err := client.Connect(ctx, clientTransport, nil)
if err != nil {
t.Fatalf("Connect() error = %v", err)
}
if got := session.InitializeResult().ProtocolVersion; got != expectedProtocolVersion {
t.Errorf("protocol version = %q, want %q", got, expectedProtocolVersion)
}
listAndCallVersion(ctx, t, session, testServerVersion)
if err := session.Close(); err != nil {
t.Fatalf("Close() error = %v", err)
}
select {
case err := <-serverDone:
if err != nil && !errors.Is(err, context.Canceled) {
t.Fatalf("server Run() error = %v", err)
}
case <-ctx.Done():
t.Fatal("server did not stop after the client session closed")
}
}
func TestStreamableHTTPStateful(t *testing.T) {
exposeAllTools(t)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
server := newMCPServer(testServerVersion)
RegisterTool(server)
httpTestServer := httptest.NewServer(newHTTPServer("", server).Handler)
defer httpTestServer.Close()
client := mcp.NewClient(&mcp.Implementation{Name: "gitea-mcp-http-test", Version: "1"}, nil)
session, err := client.Connect(ctx, &mcp.StreamableClientTransport{
Endpoint: httpTestServer.URL + "/mcp",
HTTPClient: httpTestServer.Client(),
DisableStandaloneSSE: true,
MaxRetries: -1,
}, nil)
if err != nil {
t.Fatalf("Connect() error = %v", err)
}
defer session.Close()
// Stateful Streamable HTTP cannot negotiate the sessionless 2026 protocol.
if got := session.InitializeResult().ProtocolVersion; got != expectedStatefulHTTPProtocolVersion {
t.Errorf("protocol version = %q, want %q", got, expectedStatefulHTTPProtocolVersion)
}
listAndCallVersion(ctx, t, session, testServerVersion)
response, err := httpTestServer.Client().Get(httpTestServer.URL + "/not-mcp")
if err != nil {
t.Fatalf("GET outside /mcp error = %v", err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusNotFound {
t.Errorf("GET outside /mcp status = %d, want %d", response.StatusCode, http.StatusNotFound)
}
}
// spaceReader yields an endless run of spaces, so oversized bodies can be sent
// without allocating them.
type spaceReader struct{}
func (spaceReader) Read(p []byte) (int, error) {
for index := range p {
p[index] = ' '
}
return len(p), nil
}
func TestStreamableHTTPRequestBodyLimit(t *testing.T) {
server := newMCPServer(testServerVersion)
httpTestServer := httptest.NewServer(newHTTPServer("", server).Handler)
defer httpTestServer.Close()
for _, test := range []struct {
name string
size int64
tooLarge bool
}{
{name: "above the SDK default", size: mcp.DefaultMaxRequestBodyBytes + 1},
{name: "above our own limit", size: maxRequestBodyBytes + 1, tooLarge: true},
} {
t.Run(test.name, func(t *testing.T) {
request, err := http.NewRequest(http.MethodPost, httpTestServer.URL+"/mcp", io.LimitReader(spaceReader{}, test.size))
if err != nil {
t.Fatalf("NewRequest() error = %v", err)
}
request.ContentLength = test.size
request.Header.Set("Content-Type", "application/json")
request.Header.Set("Accept", "application/json, text/event-stream")
response, err := httpTestServer.Client().Do(request)
if err != nil {
t.Fatalf("POST %d bytes error = %v", test.size, err)
}
defer response.Body.Close()
if gotTooLarge := response.StatusCode == http.StatusRequestEntityTooLarge; gotTooLarge != test.tooLarge {
t.Errorf("POST %d bytes status = %d, want %d = %v", test.size, response.StatusCode, http.StatusRequestEntityTooLarge, test.tooLarge)
}
})
}
}
type authorizationTransport struct {
base http.RoundTripper
mu sync.RWMutex
value string
}
func (t *authorizationTransport) set(value string) {
t.mu.Lock()
defer t.mu.Unlock()
t.value = value
}
func (t *authorizationTransport) RoundTrip(request *http.Request) (*http.Response, error) {
clone := request.Clone(request.Context()) // Clone already copies the header
t.mu.RLock()
value := t.value
t.mu.RUnlock()
if value != "" {
clone.Header.Set("Authorization", value)
}
return t.base.RoundTrip(clone)
}
func authContextValue(ctx context.Context, session *mcp.ClientSession) (string, error) {
result, err := session.CallTool(ctx, &mcp.CallToolParams{Name: "test_auth_context"})
if err != nil {
return "", err
}
if len(result.Content) != 1 {
return "", fmt.Errorf("content count = %d, want 1", len(result.Content))
}
content, ok := result.Content[0].(*mcp.TextContent)
if !ok {
return "", fmt.Errorf("content type = %T, want *mcp.TextContent", result.Content[0])
}
return content.Text, nil
}
func TestHTTPAuthPerRequest(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
server := newMCPServer(testServerVersion)
server.AddTool(
&mcp.Tool{
Name: "test_auth_context",
Description: "Return the request-scoped authentication token.",
InputSchema: map[string]any{"type": "object", "properties": map[string]any{}},
},
func(ctx context.Context, _ *mcp.CallToolRequest) (*mcp.CallToolResult, error) {
token, _ := ctx.Value(mcpContext.TokenContextKey).(string)
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: token}},
}, nil
},
)
httpTestServer := httptest.NewServer(newHTTPServer("", server).Handler)
defer httpTestServer.Close()
baseTransport := httpTestServer.Client().Transport
auth := &authorizationTransport{base: baseTransport}
auth.set("Bearer first-token")
baseClient := &http.Client{Transport: auth}
client := mcp.NewClient(&mcp.Implementation{Name: "gitea-mcp-auth-test", Version: "1"}, nil)
session, err := client.Connect(ctx, &mcp.StreamableClientTransport{
Endpoint: httpTestServer.URL + "/mcp",
HTTPClient: baseClient,
DisableStandaloneSSE: true,
MaxRetries: -1,
}, nil)
if err != nil {
t.Fatalf("Connect() error = %v", err)
}
defer session.Close()
for _, test := range []struct {
header string
want string
}{
{header: "Bearer first-token", want: "first-token"},
{header: "token second-token", want: "second-token"},
{header: "Basic ignored", want: ""},
} {
auth.set(test.header)
token, err := authContextValue(ctx, session)
if err != nil {
t.Fatalf("CallTool() with %q error = %v", test.header, err)
}
if token != test.want {
t.Errorf("CallTool() token = %q, want %q", token, test.want)
}
}
type authenticatedSession struct {
session *mcp.ClientSession
want string
}
concurrentSessions := make([]authenticatedSession, 0, 2)
for index, token := range []string{"parallel-one", "parallel-two"} {
transport := &authorizationTransport{base: baseTransport}
transport.set("Bearer " + token)
httpClient := &http.Client{Transport: transport}
parallelClient := mcp.NewClient(&mcp.Implementation{
Name: fmt.Sprintf("gitea-mcp-auth-parallel-%d", index),
Version: "1",
}, nil)
parallelSession, err := parallelClient.Connect(ctx, &mcp.StreamableClientTransport{
Endpoint: httpTestServer.URL + "/mcp",
HTTPClient: httpClient,
DisableStandaloneSSE: true,
MaxRetries: -1,
}, nil)
if err != nil {
t.Fatalf("parallel Connect() error = %v", err)
}
defer parallelSession.Close()
concurrentSessions = append(concurrentSessions, authenticatedSession{session: parallelSession, want: token})
}
var waitGroup sync.WaitGroup
errorsCh := make(chan error, 20)
for _, authenticated := range concurrentSessions {
for range 10 {
waitGroup.Go(func() {
got, err := authContextValue(ctx, authenticated.session)
if err != nil {
errorsCh <- err
return
}
if got != authenticated.want {
errorsCh <- fmt.Errorf("parallel token = %q, want %q", got, authenticated.want)
}
})
}
}
waitGroup.Wait()
close(errorsCh)
for err := range errorsCh {
t.Error(err)
}
}
func TestStdioCommandTransport(t *testing.T) {
if testing.Short() {
t.Skip("skipping subprocess build in short mode")
}
for name, value := range map[string]string{
"GITEA_READONLY": "true",
"GITEA_SCOPES": "user",
"GITEA_TOOLS": "get_me",
"MCP_MODE": "http",
} {
t.Setenv(name, value)
}
exposeAllTools(t)
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
binary := filepath.Join(t.TempDir(), "gitea-mcp")
build := exec.CommandContext(ctx, "go", "build", "-o", binary, "..")
if output, err := build.CombinedOutput(); err != nil {
t.Fatalf("build stdio test binary: %v\n%s", err, output)
}
client := mcp.NewClient(&mcp.Implementation{Name: "gitea-mcp-stdio-test", Version: "1"}, nil)
command := exec.CommandContext(ctx, binary, "--transport", "stdio")
command.Env = stdioCommandEnvironment()
session, err := client.Connect(ctx, &mcp.CommandTransport{
Command: command,
TerminateDuration: 2 * time.Second,
}, nil)
if err != nil {
t.Fatalf("Connect() error = %v", err)
}
defer session.Close()
if got := session.InitializeResult().ProtocolVersion; got != expectedProtocolVersion {
t.Errorf("protocol version = %q, want %q", got, expectedProtocolVersion)
}
listAndCallVersion(ctx, t, session, "Gitea MCP Server version:")
}
+53 -56
View File
@@ -13,8 +13,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("search") var Tool = tool.New("search")
@@ -27,81 +26,81 @@ const (
) )
var ( var (
SearchUsersTool = mcp.NewTool( SearchUsersTool = tool.NewDefinition(
SearchUsersToolName, SearchUsersToolName,
mcp.WithDescription("Search for Gitea users by username or full name."), "Search for Gitea users by username or full name.",
mcp.WithToolAnnotation(annotation.ReadOnly("Search users")), annotation.ReadOnly("Search users"),
mcp.WithString("query", mcp.Required()), tool.String("query", tool.Required()),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
SearOrgTeamsTool = mcp.NewTool( SearOrgTeamsTool = tool.NewDefinition(
SearchOrgTeamsToolName, SearchOrgTeamsToolName,
mcp.WithDescription("Search for teams within an organization by name, optionally including each team's description in the results."), "Search for teams within an organization by name, optionally including each team's description in the results.",
mcp.WithToolAnnotation(annotation.ReadOnly("Search organization teams")), annotation.ReadOnly("Search organization teams"),
mcp.WithString("org", mcp.Required()), tool.String("org", tool.Required()),
mcp.WithString("query", mcp.Required()), tool.String("query", tool.Required()),
mcp.WithBoolean("includeDescription"), tool.Boolean("includeDescription"),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
SearchReposTool = mcp.NewTool( SearchReposTool = tool.NewDefinition(
SearchReposToolName, SearchReposToolName,
mcp.WithDescription("Search for repositories by keyword, with filters for topic/description matching, owner, visibility, archived status, and sort order."), "Search for repositories by keyword, with filters for topic/description matching, owner, visibility, archived status, and sort order.",
mcp.WithToolAnnotation(annotation.ReadOnly("Search repositories")), annotation.ReadOnly("Search repositories"),
mcp.WithString("query", mcp.Required()), tool.String("query", tool.Required()),
mcp.WithBoolean("keywordIsTopic"), tool.Boolean("keywordIsTopic"),
mcp.WithBoolean("keywordInDescription"), tool.Boolean("keywordInDescription"),
mcp.WithNumber("ownerID"), tool.Number("ownerID"),
mcp.WithBoolean("isPrivate"), tool.Boolean("isPrivate"),
mcp.WithBoolean("isArchived"), tool.Boolean("isArchived"),
mcp.WithString("sort"), tool.String("sort"),
mcp.WithString("order"), tool.String("order"),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
SearchIssuesTool = mcp.NewTool( SearchIssuesTool = tool.NewDefinition(
SearchIssuesToolName, SearchIssuesToolName,
mcp.WithDescription("Search issues and PRs across repositories"), "Search issues and PRs across repositories",
mcp.WithToolAnnotation(annotation.ReadOnly("Search issues")), annotation.ReadOnly("Search issues"),
mcp.WithString("query", mcp.Required()), tool.String("query", tool.Required()),
mcp.WithString("state", mcp.Enum("open", "closed", "all")), tool.String("state", tool.Enum("open", "closed", "all")),
mcp.WithString("type", mcp.Enum("issues", "pulls")), tool.String("type", tool.Enum("issues", "pulls")),
mcp.WithString("labels", mcp.Description("comma-separated")), tool.String("labels", tool.Description("comma-separated")),
mcp.WithString("owner", mcp.Description("filter by owner")), tool.String("owner", tool.Description("filter by owner")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: SearchUsersTool, Tool: SearchUsersTool,
Handler: UsersFn, Handler: UsersFn,
}) })
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: SearOrgTeamsTool, Tool: SearOrgTeamsTool,
Handler: OrgTeamsFn, Handler: OrgTeamsFn,
}) })
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: SearchReposTool, Tool: SearchReposTool,
Handler: ReposFn, Handler: ReposFn,
}) })
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: SearchIssuesTool, Tool: SearchIssuesTool,
Handler: IssuesFn, Handler: IssuesFn,
}) })
} }
func UsersFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func UsersFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
keyword, err := params.GetString(req.GetArguments(), "query") keyword, err := params.GetString(args, "query")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
opt := gitea_sdk.SearchUsersOption{ opt := gitea_sdk.SearchUsersOption{
KeyWord: keyword, KeyWord: keyword,
ListOptions: gitea_sdk.ListOptions{ ListOptions: gitea_sdk.ListOptions{
@@ -120,17 +119,17 @@ func UsersFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult,
return to.TextResult(slimUserDetails(users)) return to.TextResult(slimUserDetails(users))
} }
func OrgTeamsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func OrgTeamsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
org, err := params.GetString(req.GetArguments(), "org") org, err := params.GetString(args, "org")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
query, err := params.GetString(req.GetArguments(), "query") query, err := params.GetString(args, "query")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
includeDescription, _ := req.GetArguments()["includeDescription"].(bool) includeDescription, _ := args["includeDescription"].(bool)
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
opt := gitea_sdk.SearchTeamsOptions{ opt := gitea_sdk.SearchTeamsOptions{
Query: query, Query: query,
IncludeDescription: includeDescription, IncludeDescription: includeDescription,
@@ -150,12 +149,11 @@ func OrgTeamsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResu
return to.TextResult(slimTeams(teams)) return to.TextResult(slimTeams(teams))
} }
func ReposFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func ReposFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
keyword, err := params.GetString(req.GetArguments(), "query") keyword, err := params.GetString(args, "query")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
args := req.GetArguments()
keywordIsTopic, _ := args["keywordIsTopic"].(bool) keywordIsTopic, _ := args["keywordIsTopic"].(bool)
keywordInDescription, _ := args["keywordInDescription"].(bool) keywordInDescription, _ := args["keywordInDescription"].(bool)
sort, _ := args["sort"].(string) sort, _ := args["sort"].(string)
@@ -186,8 +184,7 @@ func ReposFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult,
return to.TextResult(slim.Repos(repos)) return to.TextResult(slim.Repos(repos))
} }
func IssuesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func IssuesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
query, err := params.GetString(args, "query") query, err := params.GetString(args, "query")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+6 -4
View File
@@ -4,13 +4,13 @@ import (
"slices" "slices"
"testing" "testing"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
func TestSearchToolsRequiredFields(t *testing.T) { func TestSearchToolsRequiredFields(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
tool mcp.Tool tool *mcp.Tool
required []string required []string
}{ }{
{ {
@@ -32,9 +32,11 @@ func TestSearchToolsRequiredFields(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
inputSchema := tt.tool.InputSchema.(map[string]any)
required, _ := inputSchema["required"].([]string)
for _, field := range tt.required { for _, field := range tt.required {
if !slices.Contains(tt.tool.InputSchema.Required, field) { if !slices.Contains(required, field) {
t.Errorf("tool %s: expected %q to be required, got required=%v", tt.name, field, tt.tool.InputSchema.Required) t.Errorf("tool %s: expected %q to be required, got required=%v", tt.name, field, required)
} }
} }
}) })
+3 -1
View File
@@ -48,7 +48,9 @@ func TestSlimIssues(t *testing.T) {
} }
func TestSearchIssuesToolRequired(t *testing.T) { func TestSearchIssuesToolRequired(t *testing.T) {
if !slices.Contains(SearchIssuesTool.InputSchema.Required, "query") { inputSchema := SearchIssuesTool.InputSchema.(map[string]any)
required, _ := inputSchema["required"].([]string)
if !slices.Contains(required, "query") {
t.Error("search_issues should require query") t.Error("search_issues should require query")
} }
} }
+67 -68
View File
@@ -12,8 +12,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("timetracking") var Tool = tool.New("timetracking")
@@ -24,86 +23,86 @@ const (
) )
var ( var (
TimetrackingReadTool = mcp.NewTool( TimetrackingReadTool = tool.NewDefinition(
TimetrackingReadToolName, TimetrackingReadToolName,
mcp.WithDescription("Read time tracking: issue times, repo times, active stopwatches, your tracked times."), "Read time tracking: issue times, repo times, active stopwatches, your tracked times.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read tracked time")), annotation.ReadOnly("Read tracked time"),
mcp.WithString("method", mcp.Required(), mcp.Enum("list_issue_times", "list_repo_times", "get_my_stopwatches", "get_my_times")), tool.String("method", tool.Required(), tool.Enum("list_issue_times", "list_repo_times", "get_my_stopwatches", "get_my_times")),
mcp.WithString("owner", mcp.Description("for list_* methods")), tool.String("owner", tool.Description("for list_* methods")),
mcp.WithString("repo", mcp.Description("for list_* methods")), tool.String("repo", tool.Description("for list_* methods")),
mcp.WithNumber("issue_number", mcp.Description("for 'list_issue_times'")), tool.Number("issue_number", tool.Description("for 'list_issue_times'")),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
TimetrackingWriteTool = mcp.NewTool( TimetrackingWriteTool = tool.NewDefinition(
TimetrackingWriteToolName, TimetrackingWriteToolName,
mcp.WithDescription("Write time tracking: stopwatches and entries."), "Write time tracking: stopwatches and entries.",
mcp.WithToolAnnotation(annotation.Write("Add or manage tracked time")), annotation.Write("Add or manage tracked time"),
mcp.WithString("method", mcp.Required(), mcp.Enum("start_stopwatch", "stop_stopwatch", "delete_stopwatch", "add_time", "delete_time")), tool.String("method", tool.Required(), tool.Enum("start_stopwatch", "stop_stopwatch", "delete_stopwatch", "add_time", "delete_time")),
mcp.WithString("owner", mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Description(params.RepoDesc)), tool.String("repo", tool.Description(params.RepoDesc)),
mcp.WithNumber("issue_number"), tool.Number("issue_number"),
mcp.WithNumber("time", mcp.Description("seconds (for 'add_time')")), tool.Number("time", tool.Description("seconds (for 'add_time')")),
mcp.WithNumber("id", mcp.Description("entry ID (for 'delete_time')")), tool.Number("id", tool.Description("entry ID (for 'delete_time')")),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{Tool: TimetrackingReadTool, Handler: readFn}) Tool.RegisterRead(tool.ServerTool{Tool: TimetrackingReadTool, Handler: readFn})
Tool.RegisterWrite(server.ServerTool{Tool: TimetrackingWriteTool, Handler: writeFn}) Tool.RegisterWrite(tool.ServerTool{Tool: TimetrackingWriteTool, Handler: writeFn})
} }
func readFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func readFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "list_issue_times": case "list_issue_times":
return listTrackedTimesFn(ctx, req) return listTrackedTimesFn(ctx, args)
case "list_repo_times": case "list_repo_times":
return listRepoTimesFn(ctx, req) return listRepoTimesFn(ctx, args)
case "get_my_stopwatches": case "get_my_stopwatches":
return getMyStopwatchesFn(ctx, req) return getMyStopwatchesFn(ctx, args)
case "get_my_times": case "get_my_times":
return getMyTimesFn(ctx, req) return getMyTimesFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func writeFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func writeFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "start_stopwatch": case "start_stopwatch":
return startStopwatchFn(ctx, req) return startStopwatchFn(ctx, args)
case "stop_stopwatch": case "stop_stopwatch":
return stopStopwatchFn(ctx, req) return stopStopwatchFn(ctx, args)
case "delete_stopwatch": case "delete_stopwatch":
return deleteStopwatchFn(ctx, req) return deleteStopwatchFn(ctx, args)
case "add_time": case "add_time":
return addTrackedTimeFn(ctx, req) return addTrackedTimeFn(ctx, args)
case "delete_time": case "delete_time":
return deleteTrackedTimeFn(ctx, req) return deleteTrackedTimeFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func startStopwatchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func startStopwatchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -118,16 +117,16 @@ func startStopwatchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(fmt.Sprintf("Stopwatch started on issue %s/%s#%d", owner, repo, index)) return to.TextResult(fmt.Sprintf("Stopwatch started on issue %s/%s#%d", owner, repo, index))
} }
func stopStopwatchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func stopStopwatchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -142,16 +141,16 @@ func stopStopwatchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToo
return to.TextResult(fmt.Sprintf("Stopwatch stopped on issue %s/%s#%d - time recorded", owner, repo, index)) return to.TextResult(fmt.Sprintf("Stopwatch stopped on issue %s/%s#%d - time recorded", owner, repo, index))
} }
func deleteStopwatchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteStopwatchFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -166,7 +165,7 @@ func deleteStopwatchFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallT
return to.TextResult(fmt.Sprintf("Stopwatch deleted/cancelled on issue %s/%s#%d", owner, repo, index)) return to.TextResult(fmt.Sprintf("Stopwatch deleted/cancelled on issue %s/%s#%d", owner, repo, index))
} }
func getMyStopwatchesFn(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getMyStopwatchesFn(ctx context.Context, _ map[string]any) (*mcp.CallToolResult, error) {
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err))
@@ -181,20 +180,20 @@ func getMyStopwatchesFn(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(slimStopWatches(stopwatches)) return to.TextResult(slimStopWatches(stopwatches))
} }
func listTrackedTimesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listTrackedTimesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err))
@@ -215,21 +214,21 @@ func listTrackedTimesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Call
return to.TextResult(slimTrackedTimes(times)) return to.TextResult(slimTrackedTimes(times))
} }
func addTrackedTimeFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func addTrackedTimeFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
timeSeconds, err := params.GetIndex(req.GetArguments(), "time") timeSeconds, err := params.GetIndex(args, "time")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -246,21 +245,21 @@ func addTrackedTimeFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(slimTrackedTime(trackedTime)) return to.TextResult(slimTrackedTime(trackedTime))
} }
func deleteTrackedTimeFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteTrackedTimeFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
index, err := params.GetIndex(req.GetArguments(), "issue_number") index, err := params.GetIndex(args, "issue_number")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
id, err := params.GetIndex(req.GetArguments(), "id") id, err := params.GetIndex(args, "id")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
@@ -275,17 +274,17 @@ func deleteTrackedTimeFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Cal
return to.TextResult(fmt.Sprintf("Tracked time entry %d deleted from issue %s/%s#%d", id, owner, repo, index)) return to.TextResult(fmt.Sprintf("Tracked time entry %d deleted from issue %s/%s#%d", id, owner, repo, index))
} }
func listRepoTimesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listRepoTimesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(req.GetArguments(), "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
repo, err := params.GetString(req.GetArguments(), "repo") repo, err := params.GetString(args, "repo")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err))
@@ -305,7 +304,7 @@ func listRepoTimesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToo
return to.TextResult(slimTrackedTimes(times)) return to.TextResult(slimTrackedTimes(times))
} }
func getMyTimesFn(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getMyTimesFn(ctx context.Context, _ map[string]any) (*mcp.CallToolResult, error) {
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err))
+142
View File
@@ -0,0 +1,142 @@
package operation
import (
"encoding/json"
"slices"
"testing"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
// TestToolContract checks the properties every exposed tool must hold, rather
// than a snapshot of the current surface, so adding a tool needs no fixture
// update and a malformed schema fails here instead of panicking in AddTool.
func TestToolContract(t *testing.T) {
scopeByName := map[string]string{}
seenScopes := map[string]struct{}{}
for _, domain := range domainTools {
scope := domain.Scope()
if scope == "" {
t.Error("domainTools contains a domain with an empty scope")
}
// Tools() filters one domain by exactly one scope name, so a shared
// scope would make --scope select more than the caller asked for.
if _, duplicate := seenScopes[scope]; duplicate {
t.Errorf("domainTools contains a duplicate scope %q", scope)
}
seenScopes[scope] = struct{}{}
for _, registered := range domain.ReadTools() {
assertToolContract(t, scope, registered.Tool, true, scopeByName)
}
for _, registered := range domain.WriteTools() {
assertToolContract(t, scope, registered.Tool, false, scopeByName)
}
}
if len(scopeByName) == 0 {
t.Fatal("no tools are registered")
}
}
func assertToolContract(t *testing.T, scope string, definition *mcp.Tool, readOnly bool, scopeByName map[string]string) {
t.Helper()
t.Run(definition.Name, func(t *testing.T) {
if previous, duplicate := scopeByName[definition.Name]; duplicate {
t.Errorf("tool name is already registered in scope %q; AddTool would silently replace it", previous)
}
scopeByName[definition.Name] = scope
// Strict MCP clients reject a tools/list entry without a description.
if definition.Description == "" {
t.Error("tool has no description")
}
// A write tool registered as read stays exposed under --read-only.
if definition.Annotations == nil || definition.Annotations.ReadOnlyHint != readOnly {
t.Errorf("annotations = %+v, want readOnlyHint %v", definition.Annotations, readOnly)
}
schema := decodeJSON(t, definition.InputSchema)
if schema["type"] != "object" {
t.Fatalf("input schema type = %v, want object", schema["type"])
}
properties, ok := schema["properties"].(map[string]any)
if !ok {
t.Fatalf("input schema properties = %T, want a JSON object", schema["properties"])
}
for name, raw := range properties {
property, ok := raw.(map[string]any)
if !ok {
t.Errorf("property %q = %T, want a JSON object", name, raw)
continue
}
assertPropertyContract(t, name, property)
}
})
}
func assertPropertyContract(t *testing.T, name string, property map[string]any) {
t.Helper()
propertyType, ok := property["type"].(string)
if !ok {
t.Errorf("property %q has no type", name)
return
}
enum, hasEnum := property["enum"].([]any)
if _, declared := property["enum"]; declared && len(enum) == 0 {
t.Errorf("property %q has an empty enum", name)
}
defaultValue, hasDefault := property["default"]
if !hasDefault {
return
}
if !matchesJSONType(defaultValue, propertyType) {
t.Errorf("property %q default %#v is not a %s", name, defaultValue, propertyType)
}
if hasEnum && !slices.Contains(enum, defaultValue) {
t.Errorf("property %q default %#v is not one of its enum values %#v", name, defaultValue, enum)
}
}
func matchesJSONType(value any, propertyType string) bool {
switch propertyType {
case "string":
_, ok := value.(string)
return ok
case "number":
_, ok := value.(float64)
return ok
case "boolean":
_, ok := value.(bool)
return ok
case "array":
_, ok := value.([]any)
return ok
case "object":
_, ok := value.(map[string]any)
return ok
default:
return false
}
}
// decodeJSON round-trips through JSON so the assertions see what an MCP client
// receives rather than the Go values behind it.
func decodeJSON(t *testing.T, value any) map[string]any {
t.Helper()
encoded, err := json.Marshal(value)
if err != nil {
t.Fatalf("marshal: %v", err)
}
var decoded map[string]any
if err := json.Unmarshal(encoded, &decoded); err != nil {
t.Fatalf("decode: %v", err)
}
return decoded
}
+14 -15
View File
@@ -12,8 +12,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
gitea_sdk "gitea.dev/sdk" gitea_sdk "gitea.dev/sdk"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
const ( const (
@@ -24,27 +23,27 @@ const (
var Tool = tool.New("user") var Tool = tool.New("user")
var ( var (
GetMyUserInfoTool = mcp.NewTool( GetMyUserInfoTool = tool.NewDefinition(
GetMyUserInfoToolName, GetMyUserInfoToolName,
mcp.WithDescription("Get current user"), "Get current user",
mcp.WithToolAnnotation(annotation.ReadOnly("Get current user information")), annotation.ReadOnly("Get current user information"),
) )
GetUserOrgsTool = mcp.NewTool( GetUserOrgsTool = tool.NewDefinition(
GetUserOrgsToolName, GetUserOrgsToolName,
mcp.WithDescription("List current user's organizations"), "List current user's organizations",
mcp.WithToolAnnotation(annotation.ReadOnly("Get user organizations")), annotation.ReadOnly("Get user organizations"),
mcp.WithNumber("page", mcp.Description(params.PageDesc), mcp.DefaultNumber(1)), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1)),
mcp.WithNumber("per_page", mcp.Description(params.PaginationDesc), mcp.DefaultNumber(30)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30)),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{Tool: GetMyUserInfoTool, Handler: GetUserInfoFn}) Tool.RegisterRead(tool.ServerTool{Tool: GetMyUserInfoTool, Handler: GetUserInfoFn})
Tool.RegisterRead(server.ServerTool{Tool: GetUserOrgsTool, Handler: GetUserOrgsFn}) Tool.RegisterRead(tool.ServerTool{Tool: GetUserOrgsTool, Handler: GetUserOrgsFn})
} }
func GetUserInfoFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetUserInfoFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
client, err := gitea.ClientFromContext(ctx) client, err := gitea.ClientFromContext(ctx)
if err != nil { if err != nil {
return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err)) return to.ErrorResult(fmt.Errorf("get gitea client err: %v", err))
@@ -56,8 +55,8 @@ func GetUserInfoFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolR
return to.TextResult(slim.UserDetail(user)) return to.TextResult(slim.UserDetail(user))
} }
func GetUserOrgsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetUserOrgsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
page, pageSize := params.GetPagination(req.GetArguments(), 30) page, pageSize := params.GetPagination(args, 30)
opt := gitea_sdk.ListOrgsOptions{ opt := gitea_sdk.ListOrgsOptions{
ListOptions: gitea_sdk.ListOptions{ ListOptions: gitea_sdk.ListOptions{
+6 -7
View File
@@ -9,8 +9,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/to" "gitea.com/gitea/gitea-mcp/pkg/to"
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("version") var Tool = tool.New("version")
@@ -19,20 +18,20 @@ const (
GetGiteaMCPServerVersion = "get_gitea_mcp_server_version" GetGiteaMCPServerVersion = "get_gitea_mcp_server_version"
) )
var GetGiteaMCPServerVersionTool = mcp.NewTool( var GetGiteaMCPServerVersionTool = tool.NewDefinition(
GetGiteaMCPServerVersion, GetGiteaMCPServerVersion,
mcp.WithDescription("Get the running version of the Gitea MCP Server itself (not the Gitea instance it connects to)."), "Get the running version of the Gitea MCP Server itself (not the Gitea instance it connects to).",
mcp.WithToolAnnotation(annotation.ReadOnly("Get server version")), annotation.ReadOnly("Get server version"),
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: GetGiteaMCPServerVersionTool, Tool: GetGiteaMCPServerVersionTool,
Handler: GetGiteaMCPServerVersionFn, Handler: GetGiteaMCPServerVersionFn,
}) })
} }
func GetGiteaMCPServerVersionFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func GetGiteaMCPServerVersionFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
version := flag.Version version := flag.Version
if version == "" { if version == "" {
version = "dev" version = "dev"
+36 -43
View File
@@ -12,8 +12,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/to" "gitea.com/gitea/gitea-mcp/pkg/to"
"gitea.com/gitea/gitea-mcp/pkg/tool" "gitea.com/gitea/gitea-mcp/pkg/tool"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
var Tool = tool.New("wiki") var Tool = tool.New("wiki")
@@ -24,77 +23,76 @@ const (
) )
var ( var (
WikiReadTool = mcp.NewTool( WikiReadTool = tool.NewDefinition(
WikiReadToolName, WikiReadToolName,
mcp.WithDescription("Read wiki: list pages, get content, revision history."), "Read wiki: list pages, get content, revision history.",
mcp.WithToolAnnotation(annotation.ReadOnly("Read wiki pages")), annotation.ReadOnly("Read wiki pages"),
mcp.WithString("method", mcp.Required(), mcp.Enum("list", "get", "get_revisions")), tool.String("method", tool.Required(), tool.Enum("list", "get", "get_revisions")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("pageName", mcp.Description("for 'get'/'get_revisions'")), tool.String("pageName", tool.Description("for 'get'/'get_revisions'")),
) )
WikiWriteTool = mcp.NewTool( WikiWriteTool = tool.NewDefinition(
WikiWriteToolName, WikiWriteToolName,
mcp.WithDescription("Write wiki pages: create, update, delete."), "Write wiki pages: create, update, delete.",
mcp.WithToolAnnotation(annotation.Destructive("Create, update, or delete wiki pages")), annotation.Destructive("Create, update, or delete wiki pages"),
mcp.WithString("method", mcp.Required(), mcp.Enum("create", "update", "delete")), tool.String("method", tool.Required(), tool.Enum("create", "update", "delete")),
mcp.WithString("owner", mcp.Required(), mcp.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
mcp.WithString("repo", mcp.Required(), mcp.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
mcp.WithString("pageName", mcp.Description("for 'update'/'delete'")), tool.String("pageName", tool.Description("for 'update'/'delete'")),
mcp.WithString("title", mcp.Description("for 'create'")), tool.String("title", tool.Description("for 'create'")),
mcp.WithString("content", mcp.Description("for 'create'/'update'")), tool.String("content", tool.Description("for 'create'/'update'")),
mcp.WithString("message", mcp.Description("commit message")), tool.String("message", tool.Description("commit message")),
) )
) )
func init() { func init() {
Tool.RegisterRead(server.ServerTool{ Tool.RegisterRead(tool.ServerTool{
Tool: WikiReadTool, Tool: WikiReadTool,
Handler: wikiReadFn, Handler: wikiReadFn,
}) })
Tool.RegisterWrite(server.ServerTool{ Tool.RegisterWrite(tool.ServerTool{
Tool: WikiWriteTool, Tool: WikiWriteTool,
Handler: wikiWriteFn, Handler: wikiWriteFn,
}) })
} }
func wikiReadFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func wikiReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "list": case "list":
return listWikiPagesFn(ctx, req) return listWikiPagesFn(ctx, args)
case "get": case "get":
return getWikiPageFn(ctx, req) return getWikiPageFn(ctx, args)
case "get_revisions": case "get_revisions":
return getWikiRevisionsFn(ctx, req) return getWikiRevisionsFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func wikiWriteFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func wikiWriteFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
method, err := params.GetString(req.GetArguments(), "method") method, err := params.GetString(args, "method")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
} }
switch method { switch method {
case "create": case "create":
return createWikiPageFn(ctx, req) return createWikiPageFn(ctx, args)
case "update": case "update":
return updateWikiPageFn(ctx, req) return updateWikiPageFn(ctx, args)
case "delete": case "delete":
return deleteWikiPageFn(ctx, req) return deleteWikiPageFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
} }
func listWikiPagesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func listWikiPagesFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -113,8 +111,7 @@ func listWikiPagesFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToo
return to.TextResult(result) return to.TextResult(result)
} }
func getWikiPageFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getWikiPageFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -137,8 +134,7 @@ func getWikiPageFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolR
return to.TextResult(result) return to.TextResult(result)
} }
func getWikiRevisionsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func getWikiRevisionsFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -161,8 +157,7 @@ func getWikiRevisionsFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.Call
return to.TextResult(result) return to.TextResult(result)
} }
func createWikiPageFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func createWikiPageFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -200,8 +195,7 @@ func createWikiPageFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(result) return to.TextResult(result)
} }
func updateWikiPageFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func updateWikiPageFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
@@ -245,8 +239,7 @@ func updateWikiPageFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallTo
return to.TextResult(result) return to.TextResult(result)
} }
func deleteWikiPageFn(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { func deleteWikiPageFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
args := req.GetArguments()
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
return to.ErrorResult(err) return to.ErrorResult(err)
+1 -6
View File
@@ -11,8 +11,6 @@ import (
mcpContext "gitea.com/gitea/gitea-mcp/pkg/context" mcpContext "gitea.com/gitea/gitea-mcp/pkg/context"
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"github.com/mark3labs/mcp-go/mcp"
) )
func TestWikiWriteBase64Encoding(t *testing.T) { func TestWikiWriteBase64Encoding(t *testing.T) {
@@ -54,10 +52,7 @@ func TestWikiWriteBase64Encoding(t *testing.T) {
"title": "TestPage", "title": "TestPage",
} }
req := mcp.CallToolRequest{} result, err := wikiWriteFn(ctx, args)
req.Params.Arguments = args
result, err := wikiWriteFn(ctx, req)
if err != nil { if err != nil {
t.Fatalf("wikiWriteFn() error: %v", err) t.Fatalf("wikiWriteFn() error: %v", err)
} }
+11 -13
View File
@@ -1,18 +1,16 @@
package annotation package annotation
import "github.com/mark3labs/mcp-go/mcp" import "github.com/modelcontextprotocol/go-sdk/mcp"
func ReadOnly(title string) mcp.ToolAnnotation { func ReadOnly(title string) *mcp.ToolAnnotations {
return &mcp.ToolAnnotations{Title: title, ReadOnlyHint: true}
}
func Write(title string) *mcp.ToolAnnotations {
return &mcp.ToolAnnotations{Title: title}
}
func Destructive(title string) *mcp.ToolAnnotations {
t := true t := true
return mcp.ToolAnnotation{Title: title, ReadOnlyHint: &t} return &mcp.ToolAnnotations{Title: title, DestructiveHint: &t}
}
func Write(title string) mcp.ToolAnnotation {
f := false
return mcp.ToolAnnotation{Title: title, ReadOnlyHint: &f}
}
func Destructive(title string) mcp.ToolAnnotation {
f, t := false, true
return mcp.ToolAnnotation{Title: title, ReadOnlyHint: &f, DestructiveHint: &t}
} }
+50
View File
@@ -0,0 +1,50 @@
package annotation
import (
"encoding/json"
"maps"
"testing"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
// The hints are what clients use to decide whether a tool needs confirmation, so
// assert the encoded form: an omitted readOnlyHint reads as false either way, but
// only the explicit form survives a client that checks for the key.
func TestAnnotations(t *testing.T) {
for _, test := range []struct {
name string
annotations *mcp.ToolAnnotations
want map[string]any
}{
{
name: "ReadOnly",
annotations: ReadOnly("Read"),
want: map[string]any{"title": "Read", "readOnlyHint": true, "idempotentHint": false},
},
{
name: "Write",
annotations: Write("Write"),
want: map[string]any{"title": "Write", "readOnlyHint": false, "idempotentHint": false},
},
{
name: "Destructive",
annotations: Destructive("Delete"),
want: map[string]any{"title": "Delete", "readOnlyHint": false, "idempotentHint": false, "destructiveHint": true},
},
} {
t.Run(test.name, func(t *testing.T) {
encoded, err := json.Marshal(test.annotations)
if err != nil {
t.Fatalf("json.Marshal() error = %v", err)
}
var got map[string]any
if err := json.Unmarshal(encoded, &got); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err)
}
if !maps.Equal(got, test.want) {
t.Errorf("annotations = %s, want %v", encoded, test.want)
}
})
}
}
+4 -2
View File
@@ -7,7 +7,7 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"gitea.com/gitea/gitea-mcp/pkg/log" "gitea.com/gitea/gitea-mcp/pkg/log"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
func TextResult(v any) (*mcp.CallToolResult, error) { func TextResult(v any) (*mcp.CallToolResult, error) {
@@ -18,7 +18,9 @@ func TextResult(v any) (*mcp.CallToolResult, error) {
if flag.Debug { if flag.Debug {
log.Debugf("Text Result: %s", string(resultBytes)) log.Debugf("Text Result: %s", string(resultBytes))
} }
return mcp.NewToolResultText(string(resultBytes)), nil return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: string(resultBytes)}},
}, nil
} }
func ErrorResult(err error) (*mcp.CallToolResult, error) { func ErrorResult(err error) (*mcp.CallToolResult, error) {
+33
View File
@@ -0,0 +1,33 @@
package to
import (
"errors"
"testing"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
func TestTextResult(t *testing.T) {
result, err := TextResult(map[string]any{"name": "gitea"})
if err != nil {
t.Fatalf("TextResult() error = %v", err)
}
if len(result.Content) != 1 {
t.Fatalf("len(Content) = %d, want 1", len(result.Content))
}
content, ok := result.Content[0].(*mcp.TextContent)
if !ok {
t.Fatalf("Content[0] type = %T, want *mcp.TextContent", result.Content[0])
}
if content.Text != `{"name":"gitea"}` {
t.Errorf("Text = %q, want JSON object", content.Text)
}
}
func TestErrorResult(t *testing.T) {
want := errors.New("failed")
result, err := ErrorResult(want)
if result != nil || !errors.Is(err, want) {
t.Errorf("ErrorResult() = (%#v, %v), want (nil, %v)", result, err, want)
}
}
+106
View File
@@ -0,0 +1,106 @@
package tool
import "github.com/modelcontextprotocol/go-sdk/mcp"
// Property describes one property in a tool's input schema.
type Property struct {
name string
schema map[string]any
required bool
}
// PropertyOption configures one property in a tool's input schema.
type PropertyOption func(*Property)
// NewDefinition builds a tool definition without enabling SDK-side validation.
func NewDefinition(name, description string, annotations *mcp.ToolAnnotations, properties ...Property) *mcp.Tool {
inputProperties := make(map[string]any, len(properties))
required := make([]string, 0, len(properties))
for _, property := range properties {
inputProperties[property.name] = property.schema
if property.required {
required = append(required, property.name)
}
}
inputSchema := map[string]any{
"type": "object",
"properties": inputProperties,
}
if len(required) > 0 {
inputSchema["required"] = required
}
return &mcp.Tool{
Name: name,
Description: description,
Annotations: annotations,
InputSchema: inputSchema,
}
}
func String(name string, options ...PropertyOption) Property {
return newProperty(name, map[string]any{"type": "string"}, options...)
}
func Number(name string, options ...PropertyOption) Property {
return newProperty(name, map[string]any{"type": "number"}, options...)
}
func Boolean(name string, options ...PropertyOption) Property {
return newProperty(name, map[string]any{"type": "boolean"}, options...)
}
func Array(name string, options ...PropertyOption) Property {
return newProperty(name, map[string]any{"type": "array"}, options...)
}
func Object(name string, options ...PropertyOption) Property {
return newProperty(name, map[string]any{"type": "object", "properties": map[string]any{}}, options...)
}
func newProperty(name string, schema map[string]any, options ...PropertyOption) Property {
property := Property{name: name, schema: schema}
for _, option := range options {
option(&property)
}
return property
}
// Required marks the property as required on the parent schema. It is not a
// property-level keyword, so it never touches the emitted property schema.
func Required() PropertyOption {
return func(property *Property) {
property.required = true
}
}
func Description(description string) PropertyOption {
return func(property *Property) {
property.schema["description"] = description
}
}
func Enum(values ...string) PropertyOption {
return func(property *Property) {
property.schema["enum"] = values
}
}
func Default(value any) PropertyOption {
return func(property *Property) {
property.schema["default"] = value
}
}
func Minimum(value float64) PropertyOption {
return func(property *Property) {
property.schema["minimum"] = value
}
}
func Items(schema any) PropertyOption {
return func(property *Property) {
property.schema["items"] = schema
}
}
+71
View File
@@ -0,0 +1,71 @@
package tool
import (
"reflect"
"testing"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
func TestNewDefinition(t *testing.T) {
annotations := &mcp.ToolAnnotations{Title: "Example", ReadOnlyHint: true}
definition := NewDefinition(
"example",
"Example tool",
annotations,
String("owner", Required(), Description("repository owner"), Enum("one", "two"), Default("one")),
Number("page", Required(), Default(1), Minimum(1)),
Boolean("draft"),
Array("labels", Items(map[string]any{"type": "string"})),
Object("inputs", Description("workflow inputs")),
)
if definition.Name != "example" || definition.Description != "Example tool" {
t.Fatalf("definition = %#v", definition)
}
if definition.Annotations != annotations {
t.Fatal("NewDefinition did not preserve annotations")
}
want := map[string]any{
"type": "object",
"properties": map[string]any{
"owner": map[string]any{
"type": "string",
"description": "repository owner",
"enum": []string{"one", "two"},
"default": "one",
},
"page": map[string]any{
"type": "number",
"default": 1,
"minimum": float64(1),
},
"draft": map[string]any{"type": "boolean"},
"labels": map[string]any{
"type": "array",
"items": map[string]any{"type": "string"},
},
"inputs": map[string]any{
"type": "object",
"properties": map[string]any{},
"description": "workflow inputs",
},
},
"required": []string{"owner", "page"},
}
if !reflect.DeepEqual(definition.InputSchema, want) {
t.Errorf("InputSchema = %#v, want %#v", definition.InputSchema, want)
}
}
func TestNewDefinitionWithoutRequiredProperties(t *testing.T) {
definition := NewDefinition("empty", "", nil)
schema := definition.InputSchema.(map[string]any)
if _, ok := schema["required"]; ok {
t.Errorf("InputSchema unexpectedly contains required: %#v", schema)
}
if got := schema["properties"]; !reflect.DeepEqual(got, map[string]any{}) {
t.Errorf("properties = %#v, want empty map", got)
}
}
+105
View File
@@ -0,0 +1,105 @@
package tool
import (
"context"
"encoding/json"
"errors"
"testing"
"github.com/modelcontextprotocol/go-sdk/jsonrpc"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
func callTool(handler Handler, arguments json.RawMessage) (*mcp.CallToolResult, error) {
serverTool := ServerTool{Tool: &mcp.Tool{Name: "example"}, Handler: handler}
return serverTool.MCPHandler()(context.Background(), &mcp.CallToolRequest{
Params: &mcp.CallToolParamsRaw{Arguments: arguments},
})
}
func captureArguments(into *map[string]any) Handler {
return func(_ context.Context, arguments map[string]any) (*mcp.CallToolResult, error) {
*into = arguments
return &mcp.CallToolResult{}, nil
}
}
func TestMCPHandler(t *testing.T) {
var got map[string]any
result, err := callTool(captureArguments(&got), json.RawMessage(`{"count":2,"nested":{"enabled":true}}`))
if err != nil {
t.Fatalf("MCPHandler() error = %v", err)
}
if result == nil {
t.Fatal("MCPHandler() result is nil")
}
if got["count"] != float64(2) {
t.Errorf("count type/value = %T(%v), want float64(2)", got["count"], got["count"])
}
}
func TestMCPHandlerRejectsInvalidArguments(t *testing.T) {
called := false
handler := func(context.Context, map[string]any) (*mcp.CallToolResult, error) {
called = true
return &mcp.CallToolResult{}, nil
}
for _, arguments := range []json.RawMessage{json.RawMessage(`[]`), json.RawMessage(`"text"`), json.RawMessage(`{"broken"`)} {
_, err := callTool(handler, arguments)
assertProtocolErrorCode(t, err, jsonrpc.CodeInvalidParams)
}
if called {
t.Fatal("handler was called with invalid arguments")
}
}
// Tools without parameters are callable with an omitted or null "arguments",
// which is what clients send and what mcp-go accepted before the SDK migration.
func TestMCPHandlerAcceptsAbsentArguments(t *testing.T) {
for _, arguments := range []json.RawMessage{nil, json.RawMessage(`null`)} {
var got map[string]any
if _, err := callTool(captureArguments(&got), arguments); err != nil {
t.Fatalf("MCPHandler() with arguments %s error = %v", arguments, err)
}
if got == nil || len(got) != 0 {
t.Errorf("arguments = %#v, want an empty map", got)
}
}
}
func TestMCPHandlerConvertsErrorsAndRecoversPanics(t *testing.T) {
for _, test := range []struct {
name string
handler Handler
}{
{
name: "handler error",
handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) {
return nil, errors.New("failed")
},
},
{
name: "panic",
handler: func(context.Context, map[string]any) (*mcp.CallToolResult, error) {
panic("failed")
},
},
} {
t.Run(test.name, func(t *testing.T) {
_, err := callTool(test.handler, nil)
assertProtocolErrorCode(t, err, jsonrpc.CodeInternalError)
})
}
}
func assertProtocolErrorCode(t *testing.T, err error, want int64) {
t.Helper()
var protocolErr *jsonrpc.Error
if !errors.As(err, &protocolErr) {
t.Fatalf("error = %v, want *jsonrpc.Error", err)
}
if protocolErr.Code != want {
t.Errorf("error code = %d, want %d", protocolErr.Code, want)
}
}
+73 -12
View File
@@ -1,26 +1,38 @@
package tool package tool
import ( import (
"context"
"encoding/json"
"errors"
"fmt"
"slices" "slices"
"strings" "strings"
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"gitea.com/gitea/gitea-mcp/pkg/log" "gitea.com/gitea/gitea-mcp/pkg/log"
"github.com/mark3labs/mcp-go/server" "github.com/modelcontextprotocol/go-sdk/jsonrpc"
"github.com/modelcontextprotocol/go-sdk/mcp"
) )
type Handler func(context.Context, map[string]any) (*mcp.CallToolResult, error)
type ServerTool struct {
Tool *mcp.Tool
Handler Handler
}
type Tool struct { type Tool struct {
scope string scope string
write []server.ServerTool write []ServerTool
read []server.ServerTool read []ServerTool
} }
func New(scope string) *Tool { func New(scope string) *Tool {
return &Tool{ return &Tool{
scope: scope, scope: scope,
write: make([]server.ServerTool, 0, 100), write: make([]ServerTool, 0, 100),
read: make([]server.ServerTool, 0, 100), read: make([]ServerTool, 0, 100),
} }
} }
@@ -29,23 +41,23 @@ func (t *Tool) Scope() string {
return t.scope return t.scope
} }
func (t *Tool) RegisterWrite(s server.ServerTool) { func (t *Tool) RegisterWrite(s ServerTool) {
t.write = append(t.write, s) t.write = append(t.write, s)
} }
func (t *Tool) RegisterRead(s server.ServerTool) { func (t *Tool) RegisterRead(s ServerTool) {
t.read = append(t.read, s) t.read = append(t.read, s)
} }
// ReadTools returns the read-only tools registered on this domain, ignoring // ReadTools returns the read-only tools registered on this domain, ignoring
// the read-only and allowlist flags that Tools applies. // the read-only and allowlist flags that Tools applies.
func (t *Tool) ReadTools() []server.ServerTool { func (t *Tool) ReadTools() []ServerTool {
return t.read return t.read
} }
// WriteTools returns the write tools registered on this domain, ignoring the // WriteTools returns the write tools registered on this domain, ignoring the
// read-only and allowlist flags that Tools applies. // read-only and allowlist flags that Tools applies.
func (t *Tool) WriteTools() []server.ServerTool { func (t *Tool) WriteTools() []ServerTool {
return t.write return t.write
} }
@@ -53,8 +65,8 @@ func (t *Tool) WriteTools() []server.ServerTool {
// read-only filter and the scope/tool allowlists (union semantics: a tool is // read-only filter and the scope/tool allowlists (union semantics: a tool is
// kept if its domain's scope is in AllowedScopes OR its name is in // kept if its domain's scope is in AllowedScopes OR its name is in
// AllowedTools). With no allowlists set, all tools pass through unchanged. // AllowedTools). With no allowlists set, all tools pass through unchanged.
func (t *Tool) Tools() []server.ServerTool { func (t *Tool) Tools() []ServerTool {
all := make([]server.ServerTool, 0, len(t.write)+len(t.read)) all := make([]ServerTool, 0, len(t.write)+len(t.read))
if !flag.ReadOnly { if !flag.ReadOnly {
all = append(all, t.write...) all = append(all, t.write...)
} }
@@ -63,7 +75,7 @@ func (t *Tool) Tools() []server.ServerTool {
return all return all
} }
_, scopeAllowed := flag.AllowedScopes[t.scope] _, scopeAllowed := flag.AllowedScopes[t.scope]
filtered := make([]server.ServerTool, 0, len(all)) filtered := make([]ServerTool, 0, len(all))
for _, st := range all { for _, st := range all {
_, toolAllowed := flag.AllowedTools[st.Tool.Name] _, toolAllowed := flag.AllowedTools[st.Tool.Name]
if scopeAllowed || toolAllowed { if scopeAllowed || toolAllowed {
@@ -73,6 +85,55 @@ func (t *Tool) Tools() []server.ServerTool {
return filtered return filtered
} }
// MCPHandler adapts a project handler to the official SDK's low-level handler.
func (s ServerTool) MCPHandler() mcp.ToolHandler {
return func(ctx context.Context, req *mcp.CallToolRequest) (result *mcp.CallToolResult, err error) {
defer func() {
if recovered := recover(); recovered != nil {
panicErr := fmt.Errorf("panic recovered in %s tool handler: %v", s.Tool.Name, recovered)
log.Errorf("%s", panicErr)
err = internalError(panicErr)
}
}()
arguments, err := decodeArguments(req.Params.Arguments)
if err != nil {
return nil, err
}
result, err = s.Handler(ctx, arguments)
if err != nil {
var protocolErr *jsonrpc.Error
if errors.As(err, &protocolErr) {
return nil, err
}
// Preserve mcp-go behavior; tool-result errors are a separate change.
return nil, internalError(err)
}
return result, nil
}
}
func decodeArguments(raw json.RawMessage) (map[string]any, error) {
// An omitted and a null "arguments" both mean the tool was called without any.
if len(raw) == 0 || string(raw) == "null" {
return map[string]any{}, nil
}
var arguments map[string]any
if err := json.Unmarshal(raw, &arguments); err != nil {
return nil, &jsonrpc.Error{
Code: jsonrpc.CodeInvalidParams,
Message: fmt.Sprintf("invalid tool arguments: %v", err),
}
}
return arguments, nil
}
func internalError(err error) error {
return &jsonrpc.Error{Code: jsonrpc.CodeInternalError, Message: err.Error()}
}
// warnUnmatched logs the names present in allowlist but absent from known, // warnUnmatched logs the names present in allowlist but absent from known,
// via logUnmatched, so WarnUnmatchedAllowedTools and WarnUnmatchedAllowedScopes // via logUnmatched, so WarnUnmatchedAllowedTools and WarnUnmatchedAllowedScopes
// share the same "collect, sort, no-op when empty" logic and can't drift. // share the same "collect, sort, no-op when empty" logic and can't drift.
+4 -5
View File
@@ -6,15 +6,14 @@ import (
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"github.com/mark3labs/mcp-go/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/mark3labs/mcp-go/server"
) )
func makeTool(name string) server.ServerTool { func makeTool(name string) ServerTool {
return server.ServerTool{Tool: mcp.NewTool(name)} return ServerTool{Tool: &mcp.Tool{Name: name}}
} }
func names(sts []server.ServerTool) []string { func names(sts []ServerTool) []string {
out := make([]string, len(sts)) out := make([]string, len(sts))
for i, st := range sts { for i, st := range sts {
out[i] = st.Tool.Name out[i] = st.Tool.Name