Files
spark-mcp/internal/mcp/tools/get_application_logs.go
tao.chenandClaude 2537976722 rm: list_applications 加 limit/queue + 日志 fallback 链补 NM 直连
P0 (Python parity 缺口):
  - rm.Client.ListApps 新增 queue + limit 查询参数
  - list_applications MCP tool schema 暴露 queue (string) 和 limit
    (number, default 100), handler 透传给 RM client
  - 限流原因: 生产 RM 上无 limit 会拉回 N MB JSON 撞 MaxResponseBytes
  - limit <= 0 不发送 query 参数, 跟 RM 默认行为一致; 老的调用方
    (TestListApplications_EndToEnd 等) 不传参时行为不变
  - 4 个旧 ListApps 测试调用点跟着更新, 加 3 个新测试:
    TestListApps_WithLimit, TestListApps_WithQueue,
    TestListApps_LimitZeroNoParam

P1.3 (Python fallback 行为补全):
  - rm.Client.GetLogs 加第 4 步 fallback: 当前 3 步
    (amContainerLogs/aggregated-logs/logs) 全失败时, 调 GetApp
    解析 app.amContainerLogs 字段, GET 该 URL 直连 NodeManager
  - 新 source 名 'amContainerLogs-direct' 区分 RM-level endpoint
  - 提取 doBytesWithHosts(extraHosts...) 支持 per-call host 白名单,
    NM host 动态加进 allowedHosts, SSRF 保护不破 (URL 是 RM 响应里
    回来的, 不是 LLM 任意填的)
  - 新增 TestGetLogs_FallbackToAMDirect (成功路径) 和
    TestGetLogs_FallbackFailsWhenAMFieldMissing (amContainerLogs
    字段为空时正常返回 error)

P1.4 (零代码改动 + 文档化):
  - get_application_logs / get_application_status 工具 description
    更新, 提示 LLM raw RM JSON 里包含 amContainerLogs 字段, 可在
    aggregated-logs 全部失败时自己用 fetch_url 直连 NM
  - get_application_status 函数体不变 (本就是透传 raw JSON)

验证:
  - go build / go vet / go test 全部通过
  - 5 个新测试全 PASS, 老的 TestListApplications_EndToEnd 仍 PASS
  - 单 ListApps 调用点 (list_applications.go:81) 编译过

未提交: spark-mcp-linux-amd64 (本地 build 产物, 当前 .gitignore
没拦, 需另行决定)

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-13 15:15:21 +08:00

94 lines
2.8 KiB
Go

package tools
import (
"context"
"github.com/mark3labs/mcp-go/mcp"
"spark-mcp-go/internal/rm"
)
const GetApplicationLogsName = "get_application_logs"
const defaultTailBytes = 1 << 20
// NewGetApplicationLogsTool returns the schema for the get_application_logs MCP Tool.
func NewGetApplicationLogsTool() mcp.Tool {
return mcp.NewTool(GetApplicationLogsName,
mcp.WithDescription("Fetch YARN application logs. Tries amContainerLogs (RM endpoint, often 307 to a NodeManager) first, then aggregated-logs, then the legacy /logs endpoint, and finally the AM container's direct NodeManager URL parsed from the app response. Returns {source, text}."),
mcp.WithString("cluster_id",
mcp.Required(),
mcp.Description("ID of the configured cluster (from list_clusters)"),
),
mcp.WithString("app_id",
mcp.Required(),
mcp.Description("YARN application ID, e.g. application_1234567890_0001"),
),
mcp.WithString("container",
mcp.Description("Container ID used for the aggregated-logs fallback, e.g. container_1234567890_0001_01_000001."),
),
mcp.WithNumber("tail_bytes",
mcp.Description("Maximum bytes to return. 0 means use the server-wide response byte limit."),
),
)
}
// GetApplicationLogsHandler retrieves logs using the RM fallback chain.
func (d *Deps) GetApplicationLogsHandler(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
clusterID, err := req.RequireString("cluster_id")
if err != nil {
return errResult("get_application_logs: " + err.Error()), nil
}
appID, err := req.RequireString("app_id")
if err != nil {
return errResult("get_application_logs: " + err.Error()), nil
}
args := req.GetArguments()
container := ""
if v, ok := args["container"].(string); ok {
container = v
}
tailBytes := defaultTailBytes
if v, ok := args["tail_bytes"].(float64); ok {
tailBytes = int(v)
}
if tailBytes == 0 {
tailBytes = int(d.MaxResponseBytes)
}
callLog := startToolCall(ctx, d.Logger, GetApplicationLogsName, map[string]any{
"cluster_id": clusterID,
"app_id": appID,
"container": container,
"tail_bytes": tailBytes,
})
defer callLog.End()
cl, err := d.ClusterRepo.Get(ctx, clusterID)
if err != nil {
callLog.WithError(err)
return errResult("get_application_logs: cluster " + clusterID + ": " + err.Error()), nil
}
rmc := rm.New(d.HTTPClient, cl)
body, source, err := rmc.GetLogs(ctx, appID, container)
if err != nil {
callLog.WithError(err)
return errResult("get_application_logs: " + err.Error()), nil
}
text := string(body)
if len(text) > tailBytes {
text = truncateMiddle(text, tailBytes)
}
result := map[string]any{
"source": source,
"text": text,
}
callLog.WithResult(map[string]any{"source": source, "bytes": len(body), "returned_bytes": len(text)})
return textResult(encodeJSON(result)), nil
}