精简项目:只保留 docker_logs + health + info 三个工具

- 移除 database/middleware 等复杂模块,先聚焦核心功能
- docker_logs:封装 docker logs --tail N <container>
- 修复 go.mod 版本为 1.24 匹配 Docker 构建镜像
- Dockerfile 改用 alpine + docker-cli,避免 scratch 无 shell 问题
- docker-compose.yml 简化为单服务定义

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
yangzhaohan
2026-07-06 10:42:23 +08:00
parent fb2c3a6cae
commit 093511bbdd
23 changed files with 144 additions and 1352 deletions
-399
View File
@@ -1,399 +0,0 @@
package tool
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"log/slog"
"regexp"
"strings"
"sync"
"time"
_ "github.com/go-sql-driver/mysql"
"github.com/mark3labs/mcp-go/mcp"
"github.com/mark3labs/mcp-go/server"
_ "github.com/jackc/pgx/v5/stdlib"
)
const (
maxRows = 1000
maxSQLBytes = 4096
queryTimeout = 30 * time.Second
)
// 禁止的 SQL 关键词(大写,用于只读防护的第二层)
var forbiddenKeywords = regexp.MustCompile(
`\b(DROP|TRUNCATE|ALTER|CREATE|INSERT|UPDATE|DELETE|GRANT|REVOKE|REPLACE|LOAD|IMPORT|EXPORT)\b`,
)
// DatabaseConf 数据库连接配置
type DatabaseConf struct {
Alias string
Driver string // postgres | mysql
DSN string
}
// DatabaseTool 提供只读数据库查询工具。
type DatabaseTool struct {
configs map[string]DatabaseConf
pools map[string]*sql.DB
mu sync.RWMutex
}
func NewDatabaseTool(configs []DatabaseConf) *DatabaseTool {
cfgMap := make(map[string]DatabaseConf, len(configs))
for _, c := range configs {
cfgMap[c.Alias] = c
}
return &DatabaseTool{configs: cfgMap, pools: make(map[string]*sql.DB)}
}
func (d *DatabaseTool) Name() string { return "database" }
func (d *DatabaseTool) Description() string { return "数据库只读查询" }
func (d *DatabaseTool) Initialize(ctx context.Context) error {
for alias, cfg := range d.configs {
// 强制追加只读参数
dsn := d.enforceReadOnly(cfg)
pool, err := sql.Open(cfg.Driver, dsn)
if err != nil {
return fmt.Errorf("open %s: %w", alias, err)
}
pool.SetMaxOpenConns(5)
pool.SetMaxIdleConns(2)
pool.SetConnMaxLifetime(5 * time.Minute)
if err := pool.PingContext(ctx); err != nil {
return fmt.Errorf("ping %s: %w", alias, err)
}
d.pools[alias] = pool
slog.Info("database connected", "alias", alias, "driver", cfg.Driver)
}
return nil
}
func (d *DatabaseTool) enforceReadOnly(cfg DatabaseConf) string {
switch cfg.Driver {
case "postgres", "pgx":
if !strings.Contains(cfg.DSN, "default_transaction_read_only") {
sep := "?"
if strings.Contains(cfg.DSN, "?") {
sep = "&"
}
return cfg.DSN + sep + "default_transaction_read_only=on"
}
case "mysql":
// mysql 驱动在 DSN 中不支持这个参数,通过连接时设置
}
return cfg.DSN
}
func (d *DatabaseTool) Shutdown(_ context.Context) error {
for alias, pool := range d.pools {
if err := pool.Close(); err != nil {
slog.Error("close db pool", "alias", alias, "err", err)
}
}
return nil
}
func (d *DatabaseTool) HealthCheck(ctx context.Context) error {
for alias, pool := range d.pools {
if err := pool.PingContext(ctx); err != nil {
return fmt.Errorf("%s: %w", alias, err)
}
}
return nil
}
func (d *DatabaseTool) Register(mcpServer *server.MCPServer) error {
mcpServer.AddTool(mcp.NewTool("db_query",
mcp.WithDescription("执行只读 SQL 查询(参数化)"),
mcp.WithString("database", mcp.Required(), mcp.Description("数据库别名")),
mcp.WithString("sql", mcp.Required(), mcp.Description("SQL 查询语句")),
mcp.WithString("params", mcp.Description("JSON 数组格式的参数,如 [1, 'hello']")),
), d.handleQuery)
mcpServer.AddTool(mcp.NewTool("db_tables",
mcp.WithDescription("列出数据库中的所有表"),
mcp.WithString("database", mcp.Required(), mcp.Description("数据库别名")),
mcp.WithString("schema", mcp.Description("schema 名称(PostgreSQL),默认 public")),
), d.handleTables)
mcpServer.AddTool(mcp.NewTool("db_table_info",
mcp.WithDescription("查看表结构和索引"),
mcp.WithString("database", mcp.Required(), mcp.Description("数据库别名")),
mcp.WithString("table", mcp.Required(), mcp.Description("表名")),
), d.handleTableInfo)
mcpServer.AddTool(mcp.NewTool("db_explain",
mcp.WithDescription("EXPLAIN 分析查询计划"),
mcp.WithString("database", mcp.Required(), mcp.Description("数据库别名")),
mcp.WithString("sql", mcp.Required(), mcp.Description("要分析的 SQL")),
), d.handleExplain)
return nil
}
// --- 安全校验 ---
func (d *DatabaseTool) validateSQL(sqlStr string) error {
if len(sqlStr) > maxSQLBytes {
return fmt.Errorf("SQL too long: %d bytes (max %d)", len(sqlStr), maxSQLBytes)
}
// 禁止多语句
if strings.Contains(sqlStr, ";") {
return fmt.Errorf("multiple statements not allowed")
}
// 禁止危险关键词
upper := strings.ToUpper(sqlStr)
if loc := forbiddenKeywords.FindStringIndex(upper); loc != nil {
return fmt.Errorf("forbidden keyword detected near position %d", loc[0])
}
return nil
}
func (d *DatabaseTool) getPool(alias string) (*sql.DB, string, error) {
d.mu.RLock()
defer d.mu.RUnlock()
pool, ok := d.pools[alias]
if !ok {
return nil, "", fmt.Errorf("unknown database: %s (available: %v)", alias, d.listAliases())
}
cfg := d.configs[alias]
return pool, cfg.Driver, nil
}
func (d *DatabaseTool) listAliases() []string {
aliases := make([]string, 0, len(d.pools))
for a := range d.pools {
aliases = append(aliases, a)
}
return aliases
}
func parseParams(raw string) ([]any, error) {
if raw == "" {
return nil, nil
}
var params []any
if err := json.Unmarshal([]byte(raw), &params); err != nil {
return nil, fmt.Errorf("params must be a JSON array: %w", err)
}
return params, nil
}
// --- Handlers ---
func (d *DatabaseTool) handleQuery(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := getArgs(req)
alias, _ := args["database"].(string)
sqlStr, _ := args["sql"].(string)
rawParams, _ := args["params"].(string)
if err := d.validateSQL(sqlStr); err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
pool, driver, err := d.getPool(alias)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
params, err := parseParams(rawParams)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
ctx, cancel := context.WithTimeout(ctx, queryTimeout)
defer cancel()
// 对 MySQL 连接设置只读
if driver == "mysql" {
if _, err := pool.ExecContext(ctx, "SET SESSION TRANSACTION READ ONLY"); err == nil {
defer pool.ExecContext(context.Background(), "SET SESSION TRANSACTION READ WRITE")
}
}
sqlStr = limitQuery(sqlStr, driver)
rows, err := pool.QueryContext(ctx, sqlStr, params...)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("query error: %v", err)), nil
}
defer rows.Close()
return formatRows(rows, maxRows)
}
func (d *DatabaseTool) handleTables(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := getArgs(req)
alias, _ := args["database"].(string)
schema, _ := args["schema"].(string)
pool, driver, err := d.getPool(alias)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
var query string
switch driver {
case "postgres", "pgx":
if schema == "" {
schema = "public"
}
query = `SELECT table_name FROM information_schema.tables WHERE table_schema = $1 ORDER BY table_name`
case "mysql":
if schema == "" {
schema = ""
}
query = `SELECT table_name FROM information_schema.tables WHERE table_schema = ? ORDER BY table_name`
default:
query = `SELECT table_name FROM information_schema.tables WHERE table_schema = ? ORDER BY table_name`
}
ctx, cancel := context.WithTimeout(ctx, queryTimeout)
defer cancel()
var rows *sql.Rows
var err2 error
if schema != "" {
rows, err2 = pool.QueryContext(ctx, query, schema)
} else {
rows, err2 = pool.QueryContext(ctx, query)
}
if err2 != nil {
return mcp.NewToolResultError(err2.Error()), nil
}
defer rows.Close()
return formatRows(rows, maxRows)
}
func (d *DatabaseTool) handleTableInfo(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := getArgs(req)
alias, _ := args["database"].(string)
table, _ := args["table"].(string)
pool, driver, err := d.getPool(alias)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
// 校验表名只包含安全字符
if !regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`).MatchString(table) {
return mcp.NewToolResultError("invalid table name"), nil
}
var query string
switch driver {
case "postgres", "pgx":
query = `SELECT column_name, data_type, is_nullable, column_default FROM information_schema.columns WHERE table_name = $1 ORDER BY ordinal_position`
case "mysql":
query = `SELECT column_name, data_type, is_nullable, column_default FROM information_schema.columns WHERE table_name = ? ORDER BY ordinal_position`
default:
query = `SELECT column_name, data_type, is_nullable, column_default FROM information_schema.columns WHERE table_name = ? ORDER BY ordinal_position`
}
ctx, cancel := context.WithTimeout(ctx, queryTimeout)
defer cancel()
rows, err := pool.QueryContext(ctx, query, table)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
defer rows.Close()
return formatRows(rows, maxRows)
}
func (d *DatabaseTool) handleExplain(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := getArgs(req)
alias, _ := args["database"].(string)
sqlStr, _ := args["sql"].(string)
if err := d.validateSQL(sqlStr); err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
pool, driver, err := d.getPool(alias)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
explainPrefix := "EXPLAIN"
if driver == "mysql" {
explainPrefix = "EXPLAIN"
}
ctx, cancel := context.WithTimeout(ctx, queryTimeout)
defer cancel()
rows, err := pool.QueryContext(ctx, explainPrefix+" "+sqlStr)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
defer rows.Close()
return formatRows(rows, maxRows)
}
// --- Helpers ---
// limitQuery 追加 LIMIT 防止返回过多行(如果查询中没有 LIMIT)。
func limitQuery(sqlStr, driver string) string {
upper := strings.ToUpper(sqlStr)
if strings.Contains(upper, "LIMIT") {
return sqlStr
}
return sqlStr + fmt.Sprintf(" LIMIT %d", maxRows)
}
func formatRows(rows *sql.Rows, limit int) (*mcp.CallToolResult, error) {
columns, err := rows.Columns()
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
var result []map[string]any
count := 0
for rows.Next() {
if count >= limit {
break
}
values := make([]any, len(columns))
valuePtrs := make([]any, len(columns))
for i := range values {
valuePtrs[i] = &values[i]
}
if err := rows.Scan(valuePtrs...); err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
row := make(map[string]any, len(columns))
for i, col := range columns {
row[col] = formatValue(values[i])
}
result = append(result, row)
count++
}
data, _ := json.MarshalIndent(result, "", " ")
output := fmt.Sprintf("Rows: %d\nColumns: %s\n\n%s", count, strings.Join(columns, ", "), string(data))
return mcp.NewToolResultText(output), nil
}
func formatValue(v any) any {
switch val := v.(type) {
case []byte:
return string(val)
case time.Time:
return val.Format(time.RFC3339)
default:
return v
}
}
+41 -231
View File
@@ -3,276 +3,86 @@ package tool
import (
"bytes"
"context"
"encoding/json"
"fmt"
"log/slog"
"os/exec"
"regexp"
"strings"
"time"
"strconv"
"github.com/mark3labs/mcp-go/mcp"
"github.com/mark3labs/mcp-go/server"
)
// DockerTool 提供安全的 Docker 容器信息查询工具
// 使用 docker CLI 而非 SDK,避免跨平台编译问题,且更轻量。
type DockerTool struct {
containerRegex *regexp.Regexp
maxLogBytes int64
maxLogLines int
maxLogSince time.Duration
}
// DockerTool 提供 docker logs 查询
type DockerTool struct{}
type DockerOpts struct {
ContainerPattern string
MaxLogBytes int64
MaxLogLines int
MaxLogSince time.Duration
}
func NewDockerTool(opts DockerOpts) (*DockerTool, error) {
pattern := opts.ContainerPattern
if pattern == "" {
pattern = ".*"
}
re, err := regexp.Compile(pattern)
if err != nil {
return nil, fmt.Errorf("compile container pattern: %w", err)
}
return &DockerTool{
containerRegex: re,
maxLogBytes: opts.MaxLogBytes,
maxLogLines: opts.MaxLogLines,
maxLogSince: opts.MaxLogSince,
}, nil
func NewDockerTool() *DockerTool {
return &DockerTool{}
}
func (d *DockerTool) Name() string { return "docker" }
func (d *DockerTool) Description() string { return "Docker 容器日志和安全查询" }
func (d *DockerTool) Description() string { return "Docker 容器日志查询" }
func (d *DockerTool) Initialize(_ context.Context) error {
if _, err := exec.LookPath("docker"); err != nil {
return fmt.Errorf("docker CLI not found in PATH")
return fmt.Errorf("docker CLI not found: %w", err)
}
slog.Info("docker CLI found")
slog.Info("docker CLI ready")
return nil
}
func (d *DockerTool) Shutdown(_ context.Context) error { return nil }
func (d *DockerTool) HealthCheck(ctx context.Context) error {
return d.dockerCmd(ctx, "version").Run()
return exec.CommandContext(ctx, "docker", "version").Run()
}
func (d *DockerTool) Register(mcpServer *server.MCPServer) error {
mcpServer.AddTool(mcp.NewTool("docker_ps",
mcp.WithDescription("列出 Docker 容器"),
mcp.WithString("filter", mcp.Description("按名称过滤(支持正则)")),
mcp.WithBoolean("all", mcp.Description("是否包含已停止的容器,默认 false")),
), d.handleList)
mcpServer.AddTool(mcp.NewTool("docker_logs",
mcp.WithDescription("获取容器日志(受白名单和大小限制保护)"),
mcp.WithString("container", mcp.Required(), mcp.Description("容器名称或 ID")),
mcp.WithNumber("tail", mcp.Description("返回最后 N 行,默认 100")),
mcp.WithString("since", mcp.Description("从多久前开始,如 15m、1h,默认 15m")),
mcp.WithDescription("获取 Docker 容器日志,等价于 docker logs --tail N <container>"),
mcp.WithString("container",
mcp.Required(),
mcp.Description("容器名称或 ID"),
),
mcp.WithNumber("tail",
mcp.Description("返回最后 N 行日志,默认 100"),
),
), d.handleLogs)
mcpServer.AddTool(mcp.NewTool("docker_inspect",
mcp.WithDescription("查看容器详细信息(受白名单保护)"),
mcp.WithString("container", mcp.Required(), mcp.Description("容器名称或 ID")),
), d.handleInspect)
return nil
}
// --- 安全校验 ---
func (d *DockerTool) validateContainer(name string) error {
if !d.containerRegex.MatchString(name) {
return fmt.Errorf("container %q not in allowed pattern", name)
}
return nil
}
// --- docker CLI 封装 ---
func (d *DockerTool) dockerCmd(ctx context.Context, args ...string) *exec.Cmd {
return exec.CommandContext(ctx, "docker", args...)
}
func (d *DockerTool) dockerOutput(ctx context.Context, args ...string) ([]byte, error) {
cmd := d.dockerCmd(ctx, args...)
var stderr bytes.Buffer
cmd.Stderr = &stderr
output, err := cmd.Output()
if err != nil {
return nil, fmt.Errorf("docker %s: %v\n%s", strings.Join(args, " "), err, stderr.String())
}
return output, nil
}
// --- Handlers ---
func (d *DockerTool) handleList(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := getArgs(req)
filterName, _ := args["filter"].(string)
showAll, _ := args["all"].(bool)
dockerArgs := []string{"ps", "--format", "{{json .}}"}
if showAll {
dockerArgs = append(dockerArgs, "-a")
}
output, err := d.dockerOutput(ctx, dockerArgs...)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
type dockerPsRow struct {
Names string `json:"Names"`
Image string `json:"Image"`
Status string `json:"Status"`
CreatedAt string `json:"CreatedAt"`
Ports string `json:"Ports"`
}
var result []dockerPsRow
for _, line := range strings.Split(strings.TrimSpace(string(output)), "\n") {
if line == "" {
continue
}
var row dockerPsRow
if err := json.Unmarshal([]byte(line), &row); err != nil {
continue
}
// 应用名称过滤
if filterName != "" {
re, err := regexp.Compile(filterName)
if err != nil {
continue
}
if !re.MatchString(row.Names) {
continue
}
}
// 白名单检查
if err := d.validateContainer(row.Names); err != nil {
continue
}
result = append(result, row)
}
data, _ := json.MarshalIndent(result, "", " ")
return mcp.NewToolResultText(fmt.Sprintf("```json\n%s\n```", string(data))), nil
}
func (d *DockerTool) handleLogs(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := getArgs(req)
containerName, _ := args["container"].(string)
if err := d.validateContainer(containerName); err != nil {
return mcp.NewToolResultError(err.Error()), nil
container, _ := args["container"].(string)
if container == "" {
return mcp.NewToolResultError("container 参数必填"), nil
}
tail := "100"
if t, ok := args["tail"].(float64); ok {
tail = fmt.Sprintf("%d", int(t))
tail := 100
if t, ok := args["tail"].(float64); ok && t > 0 {
tail = int(t)
}
since := "15m"
if s, ok := args["since"].(string); ok && s != "" {
dur, err := time.ParseDuration(s)
if err == nil && d.maxLogSince > 0 && dur > d.maxLogSince {
since = d.maxLogSince.String()
} else {
since = s
}
cmd := exec.CommandContext(ctx, "docker", "logs",
"--tail", strconv.Itoa(tail),
"--timestamps",
container,
)
var stdout, stderr bytes.Buffer
cmd.Stdout = &stdout
cmd.Stderr = &stderr
if err := cmd.Run(); err != nil {
return mcp.NewToolResultError(
fmt.Sprintf("docker logs 失败: %v\n%s", err, stderr.String()),
), nil
}
output, err := d.dockerOutput(ctx, "logs", "--tail", tail, "--since", since, "--timestamps", containerName)
if err != nil {
// docker logs 返回非 0 时容器可能不存在
return mcp.NewToolResultError(fmt.Sprintf("docker logs failed: %v\nOutput: %s", err, string(output))), nil
output := stdout.String()
if output == "" {
output = "(容器没有日志输出)"
}
// 截断字节数
text := string(output)
if len(text) > int(d.maxLogBytes) {
text = text[len(text)-int(d.maxLogBytes):]
}
// 截断行数
lines := strings.Split(text, "\n")
if len(lines) > d.maxLogLines {
lines = lines[len(lines)-d.maxLogLines:]
}
return mcp.NewToolResultText(strings.Join(lines, "\n")), nil
}
func (d *DockerTool) handleInspect(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
args := getArgs(req)
containerName, _ := args["container"].(string)
if err := d.validateContainer(containerName); err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
output, err := d.dockerOutput(ctx, "inspect", containerName)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
}
// 只提取安全字段
var inspects []struct {
Name string `json:"Name"`
ID string `json:"Id"`
Image string `json:"Image"`
State struct {
Status string `json:"Status"`
Running bool `json:"Running"`
StartedAt string `json:"StartedAt"`
Pid int `json:"Pid"`
} `json:"State"`
Created string `json:"Created"`
Config struct {
Image string `json:"Image"`
Env []string `json:"Env"`
} `json:"Config"`
Mounts []struct {
Source string `json:"Source"`
Destination string `json:"Destination"`
Mode string `json:"Mode"`
} `json:"Mounts"`
}
if err := json.Unmarshal(output, &inspects); err != nil {
return mcp.NewToolResultError(fmt.Sprintf("parse inspect: %v", err)), nil
}
if len(inspects) == 0 {
return mcp.NewToolResultError("container not found"), nil
}
insp := inspects[0]
// 不暴露 Env(含密钥),只暴露安全信息
safe := map[string]any{
"name": strings.TrimPrefix(insp.Name, "/"),
"id": insp.ID[:12],
"image": insp.Config.Image,
"state": insp.State,
"created": insp.Created,
"mounts": insp.Mounts,
}
data, _ := json.MarshalIndent(safe, "", " ")
return mcp.NewToolResultText(fmt.Sprintf("```json\n%s\n```", string(data))), nil
return mcp.NewToolResultText(output), nil
}
+18 -22
View File
@@ -13,35 +13,33 @@ import (
var StartTime = time.Now()
// SystemTool 提供 health 和 info 两个基础工具。
type SystemTool struct {
version string
registry HealthChecker
}
// HealthChecker 用于 system 工具检查其他工具的连通性。
type HealthChecker interface {
HealthCheckAll(ctx context.Context) map[string]error
}
type SystemTool struct {
version string
registry HealthChecker
}
func NewSystemTool(version string, reg HealthChecker) *SystemTool {
return &SystemTool{version: version, registry: reg}
}
func (s *SystemTool) Name() string { return "system" }
func (s *SystemTool) Description() string { return "系统健康检查和信息查询" }
func (s *SystemTool) Description() string { return "服务健康检查和信息查询" }
func (s *SystemTool) Initialize(_ context.Context) error { return nil }
func (s *SystemTool) Shutdown(_ context.Context) error { return nil }
func (s *SystemTool) HealthCheck(_ context.Context) error { return nil }
func (s *SystemTool) Register(mcpServer *server.MCPServer) error {
mcpServer.AddTool(mcp.NewTool("health",
mcp.WithDescription("服务健康检查:运行时间、内存使用、连接状态"),
mcp.WithDescription("服务健康检查:运行时间、内存、连接状态"),
), s.handleHealth)
mcpServer.AddTool(mcp.NewTool("info",
mcp.WithDescription("服务信息:版本号、已配置的数据库、已加载工具"),
mcp.WithDescription("服务信息:版本号、Go 版本、已加载工具"),
), s.handleInfo)
return nil
@@ -52,11 +50,10 @@ func (s *SystemTool) handleHealth(ctx context.Context, _ mcp.CallToolRequest) (*
runtime.ReadMemStats(&m)
result := map[string]any{
"status": "ok",
"uptime": time.Since(StartTime).String(),
"goroutines": runtime.NumGoroutine(),
"heap_mb": float64(m.HeapAlloc) / 1024 / 1024,
"gc_cycles": m.NumGC,
"status": "ok",
"uptime": time.Since(StartTime).String(),
"goroutines": runtime.NumGoroutine(),
"heap_mb": float64(m.HeapAlloc) / 1024 / 1024,
}
if s.registry != nil {
@@ -69,14 +66,13 @@ func (s *SystemTool) handleHealth(ctx context.Context, _ mcp.CallToolRequest) (*
func (s *SystemTool) handleInfo(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) {
result := map[string]any{
"version": s.version,
"go_version": runtime.Version(),
"os": runtime.GOOS,
"arch": runtime.GOARCH,
"cpus": runtime.NumCPU(),
"start_time": StartTime.Format(time.RFC3339),
"version": s.version,
"go_version": runtime.Version(),
"os": runtime.GOOS,
"arch": runtime.GOARCH,
"cpus": runtime.NumCPU(),
"start_time": StartTime.Format(time.RFC3339),
}
data, _ := json.MarshalIndent(result, "", " ")
return mcp.NewToolResultText(fmt.Sprintf("```json\n%s\n```", string(data))), nil
}
+10 -21
View File
@@ -7,6 +7,16 @@ import (
"github.com/mark3labs/mcp-go/server"
)
// Tool 是所有 MCP 工具必须实现的接口。
type Tool interface {
Name() string
Description() string
Register(mcpServer *server.MCPServer) error
Initialize(ctx context.Context) error
Shutdown(ctx context.Context) error
HealthCheck(ctx context.Context) error
}
// getArgs 把 req.Params.Arguments 转换为 map[string]any。
func getArgs(req mcp.CallToolRequest) map[string]any {
if args, ok := req.Params.Arguments.(map[string]any); ok {
@@ -14,24 +24,3 @@ func getArgs(req mcp.CallToolRequest) map[string]any {
}
return map[string]any{}
}
// Tool 是所有 MCP 工具必须实现的接口。
// 每个工具模块实现此接口,然后注册到 Registry。
type Tool interface {
Name() string
// Description 返回工具的一句话描述
Description() string
// Register 向 MCP Server 注册自己的工具处理器
Register(mcpServer *server.MCPServer) error
// Initialize 建立连接、初始化资源
Initialize(ctx context.Context) error
// Shutdown 释放资源、关闭连接
Shutdown(ctx context.Context) error
// HealthCheck 返回当前工具的连接状态,nil 表示健康
HealthCheck(ctx context.Context) error
}