精简项目:只保留 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:
@@ -1,42 +0,0 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// responseWriter 捕获状态码,实现 http.ResponseWriter。
|
||||
type responseWriter struct {
|
||||
http.ResponseWriter
|
||||
status int
|
||||
}
|
||||
|
||||
func (rw *responseWriter) WriteHeader(code int) {
|
||||
rw.status = code
|
||||
rw.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
|
||||
// Audit 记录每个 HTTP 请求的结构化审计日志。
|
||||
func Audit(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
start := time.Now()
|
||||
rw := &responseWriter{ResponseWriter: w, status: 200}
|
||||
|
||||
next.ServeHTTP(rw, r)
|
||||
|
||||
subject := "anonymous"
|
||||
if sub, ok := r.Context().Value(KeySubject).(string); ok {
|
||||
subject = sub
|
||||
}
|
||||
|
||||
slog.Info("request",
|
||||
"method", r.Method,
|
||||
"path", r.URL.Path,
|
||||
"status", rw.status,
|
||||
"duration_ms", time.Since(start).Milliseconds(),
|
||||
"subject", subject,
|
||||
"remote_addr", r.RemoteAddr,
|
||||
)
|
||||
})
|
||||
}
|
||||
@@ -1,94 +0,0 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
type contextKey string
|
||||
|
||||
const (
|
||||
KeySubject contextKey = "subject"
|
||||
KeyScopes contextKey = "scopes"
|
||||
)
|
||||
|
||||
// Auth 验证 Bearer Token(API Key 或 JWT)。
|
||||
type Auth struct {
|
||||
apiKeys map[string]string // key → description
|
||||
jwtSecret []byte
|
||||
}
|
||||
|
||||
func NewAuth(apiKeys map[string]string, jwtSecret string) *Auth {
|
||||
return &Auth{
|
||||
apiKeys: apiKeys,
|
||||
jwtSecret: []byte(jwtSecret),
|
||||
}
|
||||
}
|
||||
|
||||
// Middleware 从 HTTP Header 中提取并验证 Token。
|
||||
func (a *Auth) Middleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
token := extractBearer(r)
|
||||
if token == "" {
|
||||
http.Error(w, "missing Bearer token", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
// 尝试 API Key 验证
|
||||
if desc, ok := a.apiKeys[token]; ok {
|
||||
ctx := context.WithValue(r.Context(), KeySubject, fmt.Sprintf("apikey:%s", desc))
|
||||
ctx = context.WithValue(ctx, KeyScopes, []string{"*"})
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
return
|
||||
}
|
||||
|
||||
// 尝试 JWT 验证
|
||||
if a.jwtSecret != nil {
|
||||
claims, err := a.parseJWT(token)
|
||||
if err == nil {
|
||||
ctx := context.WithValue(r.Context(), KeySubject, claims.Subject)
|
||||
ctx = context.WithValue(ctx, KeyScopes, claims.Scopes)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
http.Error(w, "invalid token", http.StatusUnauthorized)
|
||||
})
|
||||
}
|
||||
|
||||
type customClaims struct {
|
||||
jwt.RegisteredClaims
|
||||
Scopes []string `json:"scopes"`
|
||||
}
|
||||
|
||||
func (a *Auth) parseJWT(tokenStr string) (*customClaims, error) {
|
||||
token, err := jwt.ParseWithClaims(tokenStr, &customClaims{},
|
||||
func(t *jwt.Token) (any, error) {
|
||||
return a.jwtSecret, nil
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
claims, ok := token.Claims.(*customClaims)
|
||||
if !ok || !token.Valid {
|
||||
return nil, fmt.Errorf("invalid claims")
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
func extractBearer(r *http.Request) string {
|
||||
auth := r.Header.Get("Authorization")
|
||||
if auth == "" {
|
||||
return ""
|
||||
}
|
||||
if !strings.HasPrefix(auth, "Bearer ") {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimPrefix(auth, "Bearer ")
|
||||
}
|
||||
@@ -1,53 +0,0 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"sync"
|
||||
|
||||
"golang.org/x/time/rate"
|
||||
)
|
||||
|
||||
// RateLimiter 基于调用方标识的令牌桶限流。
|
||||
type RateLimiter struct {
|
||||
limiters map[string]*rate.Limiter
|
||||
rate rate.Limit
|
||||
burst int
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewRateLimiter(r rate.Limit, burst int) *RateLimiter {
|
||||
return &RateLimiter{
|
||||
limiters: make(map[string]*rate.Limiter),
|
||||
rate: r,
|
||||
burst: burst,
|
||||
}
|
||||
}
|
||||
|
||||
func (rl *RateLimiter) Middleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
subject := "anonymous"
|
||||
if sub, ok := r.Context().Value(KeySubject).(string); ok {
|
||||
subject = sub
|
||||
}
|
||||
|
||||
limiter := rl.getLimiter(subject)
|
||||
if !limiter.Allow() {
|
||||
http.Error(w, "rate limit exceeded", http.StatusTooManyRequests)
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func (rl *RateLimiter) getLimiter(key string) *rate.Limiter {
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
|
||||
limiter, ok := rl.limiters[key]
|
||||
if !ok {
|
||||
limiter = rate.NewLimiter(rl.rate, rl.burst)
|
||||
rl.limiters[key] = limiter
|
||||
}
|
||||
return limiter
|
||||
}
|
||||
@@ -1,17 +0,0 @@
|
||||
package server
|
||||
|
||||
// Config 是 server 包所需的配置子集,避免循环依赖。
|
||||
// 实际值由 main.go 从 config.Config 映射过来。
|
||||
type Config struct {
|
||||
Server struct {
|
||||
Transport string
|
||||
Addr string
|
||||
}
|
||||
Docker struct {
|
||||
Enabled bool
|
||||
}
|
||||
Databases map[string]struct {
|
||||
Driver string
|
||||
DSN string
|
||||
}
|
||||
}
|
||||
@@ -19,18 +19,16 @@ func NewRegistry() *Registry {
|
||||
return &Registry{tools: make(map[string]tool.Tool)}
|
||||
}
|
||||
|
||||
// Register 注册一个工具。如果工具名重复则报错。
|
||||
func (r *Registry) Register(t tool.Tool) error {
|
||||
name := t.Name()
|
||||
if _, exists := r.tools[name]; exists {
|
||||
return fmt.Errorf("tool %q already registered", name)
|
||||
}
|
||||
r.tools[name] = t
|
||||
slog.Info("tool registered", "name", name, "desc", t.Description())
|
||||
slog.Info("tool registered", "name", name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// InitializeAll 调用所有工具的 Initialize,任一失败则终止。
|
||||
func (r *Registry) InitializeAll(ctx context.Context) error {
|
||||
for name, t := range r.tools {
|
||||
slog.Info("initializing tool", "name", name)
|
||||
@@ -41,18 +39,15 @@ func (r *Registry) InitializeAll(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// RegisterAll 将所有工具注册到 MCP Server。
|
||||
func (r *Registry) RegisterAll(mcpServer *server.MCPServer) error {
|
||||
for name, t := range r.tools {
|
||||
if err := t.Register(mcpServer); err != nil {
|
||||
return fmt.Errorf("register tool %q: %w", name, err)
|
||||
}
|
||||
slog.Info("tool handlers registered", "name", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ShutdownAll 调用所有工具的 Shutdown,收集所有错误。
|
||||
func (r *Registry) ShutdownAll(ctx context.Context) []error {
|
||||
var errs []error
|
||||
for name, t := range r.tools {
|
||||
@@ -64,7 +59,6 @@ func (r *Registry) ShutdownAll(ctx context.Context) []error {
|
||||
return errs
|
||||
}
|
||||
|
||||
// HealthCheckAll 检查所有工具的连通性,用于 system health 工具。
|
||||
func (r *Registry) HealthCheckAll(ctx context.Context) map[string]error {
|
||||
result := make(map[string]error, len(r.tools))
|
||||
for name, t := range r.tools {
|
||||
@@ -72,12 +66,3 @@ func (r *Registry) HealthCheckAll(ctx context.Context) map[string]error {
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// List 返回所有已注册工具的名称。
|
||||
func (r *Registry) List() []string {
|
||||
names := make([]string, 0, len(r.tools))
|
||||
for name := range r.tools {
|
||||
names = append(names, name)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
+10
-53
@@ -4,7 +4,6 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
@@ -16,12 +15,11 @@ import (
|
||||
|
||||
const version = "0.1.0"
|
||||
|
||||
// Run 启动 MCP Server,处理信号优雅关闭。
|
||||
func Run(cfg *Config, tools ...tool.Tool) error {
|
||||
// Run 启动 MCP Server。
|
||||
func Run(transport string, tools ...tool.Tool) error {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
// 监听终止信号
|
||||
sigCh := make(chan os.Signal, 1)
|
||||
signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
|
||||
go func() {
|
||||
@@ -30,7 +28,6 @@ func Run(cfg *Config, tools ...tool.Tool) error {
|
||||
cancel()
|
||||
}()
|
||||
|
||||
// 创建注册中心
|
||||
registry := NewRegistry()
|
||||
for _, t := range tools {
|
||||
if err := registry.Register(t); err != nil {
|
||||
@@ -38,13 +35,11 @@ func Run(cfg *Config, tools ...tool.Tool) error {
|
||||
}
|
||||
}
|
||||
|
||||
// 系统工具总是最后注册,确保它能看到所有工具
|
||||
sysTool := tool.NewSystemTool(version, registry)
|
||||
if err := registry.Register(sysTool); err != nil {
|
||||
return fmt.Errorf("register system tool: %w", err)
|
||||
return fmt.Errorf("register system: %w", err)
|
||||
}
|
||||
|
||||
// 初始化所有工具
|
||||
if err := registry.InitializeAll(ctx); err != nil {
|
||||
return fmt.Errorf("initialize: %w", err)
|
||||
}
|
||||
@@ -54,58 +49,20 @@ func Run(cfg *Config, tools ...tool.Tool) error {
|
||||
}
|
||||
}()
|
||||
|
||||
// 创建 MCP Server
|
||||
mcpServer := server.NewMCPServer(
|
||||
"ops-mcp",
|
||||
version,
|
||||
server.WithLogging(),
|
||||
)
|
||||
mcpServer := server.NewMCPServer("ops-mcp", version, server.WithLogging())
|
||||
|
||||
// 注册所有工具处理器
|
||||
if err := registry.RegisterAll(mcpServer); err != nil {
|
||||
return fmt.Errorf("register all: %w", err)
|
||||
}
|
||||
|
||||
// 选择传输模式
|
||||
switch cfg.Server.Transport {
|
||||
switch transport {
|
||||
case "stdio":
|
||||
return runStdio(ctx, mcpServer)
|
||||
slog.Info("starting MCP server", "transport", "stdio")
|
||||
return server.ServeStdio(mcpServer)
|
||||
case "sse":
|
||||
return runSSE(ctx, mcpServer, cfg)
|
||||
slog.Info("starting MCP server", "transport", "sse", "addr", ":8080")
|
||||
return server.NewSSEServer(mcpServer).Start(":8080")
|
||||
default:
|
||||
return fmt.Errorf("unknown transport: %s (expect stdio or sse)", cfg.Server.Transport)
|
||||
return fmt.Errorf("unknown transport: %s", transport)
|
||||
}
|
||||
}
|
||||
|
||||
func runStdio(ctx context.Context, mcpServer *server.MCPServer) error {
|
||||
slog.Info("starting MCP server", "transport", "stdio")
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
slog.Info("shutting down stdio server")
|
||||
}()
|
||||
return server.ServeStdio(mcpServer)
|
||||
}
|
||||
|
||||
func runSSE(ctx context.Context, mcpServer *server.MCPServer, cfg *Config) error {
|
||||
sseServer := server.NewSSEServer(mcpServer)
|
||||
|
||||
addr := cfg.Server.Addr
|
||||
if addr == "" {
|
||||
addr = ":8080"
|
||||
}
|
||||
|
||||
slog.Info("starting MCP server", "transport", "sse", "addr", addr)
|
||||
|
||||
httpServer := &http.Server{Addr: addr, Handler: sseServer}
|
||||
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
slog.Info("shutting down HTTP server")
|
||||
httpServer.Shutdown(context.Background())
|
||||
}()
|
||||
|
||||
if err := httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -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), ¶ms); 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
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user