refactor(mcp): dedupe log fetch/collect between log tools (#4835)

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Amir Raminfar
2026-07-15 07:56:42 -07:00
committed by GitHub
parent a0704d5c7b
commit 4e7b6eea7b
2 changed files with 183 additions and 211 deletions
+147 -211
View File
@@ -169,95 +169,119 @@ func (s *Server) handleListContainers(ctx context.Context, _ *mcp.CallToolReques
}, nil, nil
}
func (s *Server) handleGetContainerLogs(ctx context.Context, _ *mcp.CallToolRequest, params *getContainerLogsParams) (*mcp.CallToolResult, any, error) {
if params.Host == "" || params.ContainerID == "" {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: "host and container_id are required"}},
IsError: true,
}, nil, nil
}
const maxLogSize = 1024 * 1024 // 1MB limit
containerSvc, err := s.hostService.FindContainer(params.Host, params.ContainerID, s.labels)
if err != nil {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: fmt.Sprintf("container not found: %v", err)}},
IsError: true,
}, nil, nil
}
// mcpLogEntry is the JSON shape returned by the log tools.
type mcpLogEntry struct {
Timestamp string `json:"timestamp"`
Level string `json:"level,omitempty"`
Stream string `json:"stream,omitempty"`
Type string `json:"type"`
Message any `json:"message"`
}
stream := ""
if params.Stream != nil {
stream = *params.Stream
func errorResult(text string) *mcp.CallToolResult {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: text}},
IsError: true,
}
var stdType container.StdType
switch stream {
}
func textResult(text string) *mcp.CallToolResult {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: text}},
}
}
// parseStream resolves the optional stream parameter into a StdType. It returns
// a non-nil error result when the value is invalid.
func parseStream(stream *string) (container.StdType, *mcp.CallToolResult) {
value := ""
if stream != nil {
value = *stream
}
switch value {
case "", "all":
stdType = container.STDALL
return container.STDALL, nil
case "stdout":
stdType = container.STDOUT
return container.STDOUT, nil
case "stderr":
stdType = container.STDERR
return container.STDERR, nil
default:
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: fmt.Sprintf("invalid stream %q: must be stdout, stderr, or all", stream)}},
IsError: true,
}, nil, nil
return container.STDALL, errorResult(fmt.Sprintf("invalid stream %q: must be stdout, stderr, or all", value))
}
}
// fetchLogs finds the container and returns its log events for the last
// sinceMinutes (defaulting to 5). On a user-facing failure it returns a non-nil
// error result; otherwise the caller must call cancel when done.
func (s *Server) fetchLogs(ctx context.Context, host, containerID string, stream *string, sinceMinutes *int) (<-chan *container.LogEvent, context.CancelFunc, *mcp.CallToolResult) {
containerSvc, err := s.hostService.FindContainer(host, containerID, s.labels)
if err != nil {
return nil, nil, errorResult(fmt.Sprintf("container not found: %v", err))
}
sinceMinutes := 5
if params.SinceMinutes != nil && *params.SinceMinutes > 0 {
sinceMinutes = *params.SinceMinutes
stdType, errResult := parseStream(stream)
if errResult != nil {
return nil, nil, errResult
}
minutes := 5
if sinceMinutes != nil && *sinceMinutes > 0 {
minutes = *sinceMinutes
}
logCtx, cancel := context.WithCancel(ctx)
defer cancel()
since := time.Now().Add(-time.Duration(sinceMinutes) * time.Minute)
since := time.Now().Add(-time.Duration(minutes) * time.Minute)
events, err := containerSvc.LogsBetweenDates(logCtx, since, time.Now(), stdType)
if err != nil {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: fmt.Sprintf("failed to read logs: %v", err)}},
IsError: true,
}, nil, nil
cancel()
return nil, nil, errorResult(fmt.Sprintf("failed to read logs: %v", err))
}
type logEntry struct {
Timestamp string `json:"timestamp"`
Level string `json:"level,omitempty"`
Stream string `json:"stream,omitempty"`
Type string `json:"type"`
Message any `json:"message"`
}
return events, cancel, nil
}
var entries []logEntry
totalSize := 0
const maxSize = 1024 * 1024 // 1MB limit
for event := range events {
var msg any
switch event.Type {
case container.LogTypeGroup:
if fragments, ok := event.Message.([]container.LogFragment); ok {
lines := make([]string, len(fragments))
for i, f := range fragments {
lines[i] = f.Message
}
msg = lines
} else {
msg = event.RawMessage
// eventMessage extracts the JSON-encodable message payload for a log event.
func eventMessage(event *container.LogEvent) any {
switch event.Type {
case container.LogTypeGroup:
if fragments, ok := event.Message.([]container.LogFragment); ok {
lines := make([]string, len(fragments))
for i, f := range fragments {
lines[i] = f.Message
}
case container.LogTypeComplex:
msg = event.Message
default:
msg = event.RawMessage
return lines
}
return event.RawMessage
case container.LogTypeComplex:
return event.Message
default:
return event.RawMessage
}
}
entry := logEntry{
Timestamp: time.UnixMilli(event.Timestamp).UTC().Format(time.RFC3339Nano),
Level: event.Level,
Stream: event.Stream,
Type: string(event.Type),
Message: msg,
func newLogEntry(event *container.LogEvent) mcpLogEntry {
return mcpLogEntry{
Timestamp: time.UnixMilli(event.Timestamp).UTC().Format(time.RFC3339Nano),
Level: event.Level,
Stream: event.Stream,
Type: string(event.Type),
Message: eventMessage(event),
}
}
// collectLogEntries drains events into JSON-encodable entries. keep, when
// non-nil, filters which entries are included. Collection stops once the encoded
// size would exceed maxLogSize, in which case truncated is true. scanned counts
// every event pulled from the channel, matched or not.
func collectLogEntries(events <-chan *container.LogEvent, keep func(mcpLogEntry) bool) (entries []mcpLogEntry, scanned int, truncated bool) {
totalSize := 0
for event := range events {
scanned++
entry := newLogEntry(event)
if keep != nil && !keep(entry) {
continue
}
line, err := json.Marshal(entry)
@@ -266,182 +290,94 @@ func (s *Server) handleGetContainerLogs(ctx context.Context, _ *mcp.CallToolRequ
}
totalSize += len(line) + 1
if totalSize > maxSize {
if totalSize > maxLogSize {
truncated = true
break
}
entries = append(entries, entry)
}
return entries, scanned, truncated
}
if len(entries) == 0 {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: "(no logs in the specified time range)"}},
}, nil, nil
}
// encodeLogEntries renders entries as newline-delimited JSON.
func encodeLogEntries(entries []mcpLogEntry) (string, error) {
var sb strings.Builder
encoder := json.NewEncoder(&sb)
for _, entry := range entries {
if err := encoder.Encode(entry); err != nil {
return nil, nil, fmt.Errorf("failed to encode log entry: %w", err)
return "", fmt.Errorf("failed to encode log entry: %w", err)
}
}
return strings.TrimRight(sb.String(), "\n"), nil
}
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: strings.TrimRight(sb.String(), "\n")}},
}, nil, nil
func (s *Server) handleGetContainerLogs(ctx context.Context, _ *mcp.CallToolRequest, params *getContainerLogsParams) (*mcp.CallToolResult, any, error) {
if params.Host == "" || params.ContainerID == "" {
return errorResult("host and container_id are required"), nil, nil
}
events, cancel, errResult := s.fetchLogs(ctx, params.Host, params.ContainerID, params.Stream, params.SinceMinutes)
if errResult != nil {
return errResult, nil, nil
}
defer cancel()
entries, _, _ := collectLogEntries(events, nil)
if len(entries) == 0 {
return textResult("(no logs in the specified time range)"), nil, nil
}
text, err := encodeLogEntries(entries)
if err != nil {
return nil, nil, err
}
return textResult(text), nil, nil
}
func (s *Server) handleSearchContainerLogs(ctx context.Context, _ *mcp.CallToolRequest, params *searchContainerLogsParams) (*mcp.CallToolResult, any, error) {
if params.Host == "" || params.ContainerID == "" || params.Query == "" {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: "host, container_id, and query are required"}},
IsError: true,
}, nil, nil
}
containerSvc, err := s.hostService.FindContainer(params.Host, params.ContainerID, s.labels)
if err != nil {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: fmt.Sprintf("container not found: %v", err)}},
IsError: true,
}, nil, nil
}
stream := ""
if params.Stream != nil {
stream = *params.Stream
}
var stdType container.StdType
switch stream {
case "", "all":
stdType = container.STDALL
case "stdout":
stdType = container.STDOUT
case "stderr":
stdType = container.STDERR
default:
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: fmt.Sprintf("invalid stream %q: must be stdout, stderr, or all", stream)}},
IsError: true,
}, nil, nil
}
sinceMinutes := 5
if params.SinceMinutes != nil && *params.SinceMinutes > 0 {
sinceMinutes = *params.SinceMinutes
}
caseSensitive := false
if params.CaseSensitive != nil {
caseSensitive = *params.CaseSensitive
return errorResult("host, container_id, and query are required"), nil, nil
}
caseSensitive := params.CaseSensitive != nil && *params.CaseSensitive
query := params.Query
if !caseSensitive {
query = strings.ToLower(query)
}
logCtx, cancel := context.WithCancel(ctx)
events, cancel, errResult := s.fetchLogs(ctx, params.Host, params.ContainerID, params.Stream, params.SinceMinutes)
if errResult != nil {
return errResult, nil, nil
}
defer cancel()
since := time.Now().Add(-time.Duration(sinceMinutes) * time.Minute)
events, err := containerSvc.LogsBetweenDates(logCtx, since, time.Now(), stdType)
if err != nil {
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: fmt.Sprintf("failed to read logs: %v", err)}},
IsError: true,
}, nil, nil
}
type logEntry struct {
Timestamp string `json:"timestamp"`
Level string `json:"level,omitempty"`
Stream string `json:"stream,omitempty"`
Type string `json:"type"`
Message any `json:"message"`
}
var entries []logEntry
totalSize := 0
const maxSize = 1024 * 1024 // 1MB limit
matchCount := 0
scanned := 0
for event := range events {
scanned++
var msg any
switch event.Type {
case container.LogTypeGroup:
if fragments, ok := event.Message.([]container.LogFragment); ok {
lines := make([]string, len(fragments))
for i, f := range fragments {
lines[i] = f.Message
}
msg = lines
} else {
msg = event.RawMessage
}
case container.LogTypeComplex:
msg = event.Message
default:
msg = event.RawMessage
}
searchStr := messageToSearchString(msg)
if caseSensitive {
if !strings.Contains(searchStr, query) {
continue
}
} else {
if !strings.Contains(strings.ToLower(searchStr), query) {
continue
}
}
matchCount++
entry := logEntry{
Timestamp: time.UnixMilli(event.Timestamp).UTC().Format(time.RFC3339Nano),
Level: event.Level,
Stream: event.Stream,
Type: string(event.Type),
Message: msg,
}
line, err := json.Marshal(entry)
if err != nil {
continue
}
totalSize += len(line) + 1
if totalSize > maxSize {
break
}
entries = append(entries, entry)
keep := func(entry mcpLogEntry) bool {
haystack := messageToSearchString(entry.Message)
if !caseSensitive {
haystack = strings.ToLower(haystack)
}
return strings.Contains(haystack, query)
}
entries, scanned, truncated := collectLogEntries(events, keep)
if len(entries) == 0 {
summary := fmt.Sprintf("(no matches for %q in %d log entries scanned)", params.Query, scanned)
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: summary}},
}, nil, nil
return textResult(fmt.Sprintf("(no matches for %q in %d log entries scanned)", params.Query, scanned)), nil, nil
}
body, err := encodeLogEntries(entries)
if err != nil {
return nil, nil, err
}
var sb strings.Builder
fmt.Fprintf(&sb, "Found %d matches for %q (scanned %d entries):\n", matchCount, params.Query, scanned)
encoder := json.NewEncoder(&sb)
for _, entry := range entries {
if err := encoder.Encode(entry); err != nil {
return nil, nil, fmt.Errorf("failed to encode log entry: %w", err)
}
fmt.Fprintf(&sb, "Found %d matches for %q (scanned %d entries):\n", len(entries), params.Query, scanned)
sb.WriteString(body)
if truncated {
sb.WriteString("\n(results truncated at 1MB; narrow your query or time range to see more)")
}
return &mcp.CallToolResult{
Content: []mcp.Content{&mcp.TextContent{Text: strings.TrimRight(sb.String(), "\n")}},
}, nil, nil
return textResult(sb.String()), nil, nil
}
func messageToSearchString(msg any) string {
+36
View File
@@ -4,6 +4,7 @@ import (
"context"
"fmt"
"io"
"strings"
"testing"
"time"
@@ -500,3 +501,38 @@ func TestSearchContainerLogsRequiredParams(t *testing.T) {
require.NoError(t, err)
assert.True(t, result.IsError)
}
func TestSearchContainerLogsTruncates(t *testing.T) {
now := time.Now()
big := strings.Repeat("match ", 200) // ~1.2KB per entry
events := make([]*container.LogEvent, 2000)
for i := range events {
events[i] = &container.LogEvent{Timestamp: now.UnixMilli(), Level: "info", Stream: "stdout", Type: container.LogTypeSingle, RawMessage: big}
}
svc := &mockHostService{
containers: []container.Container{{ID: "abc123", Name: "web", Host: "local"}},
logEvents: events,
}
s := NewServer(svc, nil, "test")
ctx := context.Background()
ct, st := mcp.NewInMemoryTransports()
_, err := s.mcpServer.Connect(ctx, st, nil)
require.NoError(t, err)
client := mcp.NewClient(&mcp.Implementation{Name: "test-client", Version: "v0.0.1"}, nil)
session, err := client.Connect(ctx, ct, nil)
require.NoError(t, err)
defer session.Close()
result, err := session.CallTool(ctx, &mcp.CallToolParams{
Name: "search_container_logs",
Arguments: map[string]any{"host": "local", "container_id": "abc123", "query": "match"},
})
require.NoError(t, err)
assert.False(t, result.IsError)
text := result.Content[0].(*mcp.TextContent).Text
assert.Contains(t, text, "results truncated")
}