Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion pkg/audit/auditor.go
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,15 @@ func NewAuditorWithTransport(config *Config, transportType string) (*Auditor, er
// Close closes the underlying log writer if it implements io.Closer.
// This should be called when the auditor is no longer needed to properly release resources.
func (a *Auditor) Close() error {
if closer, ok := a.logWriter.(io.Closer); ok {
return closeLogWriter(a.logWriter)
}

func closeLogWriter(logWriter io.Writer) error {
if logWriter == os.Stdout || logWriter == os.Stderr {
return nil
}

if closer, ok := logWriter.(io.Closer); ok {
return closer.Close()
}
return nil
Expand Down
9 changes: 9 additions & 0 deletions pkg/audit/workflow_auditor.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"time"

Expand All @@ -21,6 +22,7 @@ type WorkflowAuditor struct {
auditLogger *slog.Logger
config *Config
component string
logWriter io.Writer
}

// NewWorkflowAuditor creates a new workflow auditor.
Expand All @@ -45,9 +47,16 @@ func NewWorkflowAuditor(config *Config) (*WorkflowAuditor, error) {
auditLogger: NewAuditLogger(logWriter),
config: config,
component: component,
logWriter: logWriter,
}, nil
}

// Close closes the underlying log writer if it owns a closeable resource.
// This should be called when the workflow auditor is no longer needed.
func (w *WorkflowAuditor) Close() error {
return closeLogWriter(w.logWriter)
}

// LogWorkflowStarted logs the start of workflow execution.
func (w *WorkflowAuditor) LogWorkflowStarted(
ctx context.Context,
Expand Down
63 changes: 63 additions & 0 deletions pkg/audit/workflow_auditor_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -23,11 +23,25 @@ type testLogWriter struct {
logs []string
}

type closeTrackingWriter struct {
closed bool
closeErr error
}

func (w *testLogWriter) Write(p []byte) (n int, err error) {
w.logs = append(w.logs, string(p))
return len(p), nil
}

func (*closeTrackingWriter) Write(p []byte) (n int, err error) {
return len(p), nil
}

func (w *closeTrackingWriter) Close() error {
w.closed = true
return w.closeErr
}

func (w *testLogWriter) getLastLog() string {
if len(w.logs) == 0 {
return ""
Expand All @@ -52,6 +66,7 @@ func createTestAuditor(t *testing.T, config *Config) (*WorkflowAuditor, *testLog
auditLogger: NewAuditLogger(writer),
config: config,
component: "vmcp-composer",
logWriter: writer,
}

return auditor, writer
Expand Down Expand Up @@ -122,6 +137,54 @@ func TestNewWorkflowAuditor(t *testing.T) {
}
}

func TestWorkflowAuditor_Close(t *testing.T) {
t.Parallel()

t.Run("closes retained file writer", func(t *testing.T) {
t.Parallel()

logFilePath := t.TempDir() + "/workflow-audit.log"
auditor, err := NewWorkflowAuditor(&Config{LogFile: logFilePath})
require.NoError(t, err)

_, ok := auditor.logWriter.(interface{ Close() error })
require.True(t, ok, "file-backed workflow auditor should retain a closeable writer")

require.NoError(t, auditor.Close())
})

t.Run("does not close stdout", func(t *testing.T) {
t.Parallel()

auditor, err := NewWorkflowAuditor(&Config{})
require.NoError(t, err)

require.NoError(t, auditor.Close())
assert.Same(t, os.Stdout, auditor.logWriter)
})

t.Run("propagates close errors", func(t *testing.T) {
t.Parallel()

closeErr := errors.New("close failed")
writer := &closeTrackingWriter{closeErr: closeErr}
auditor := &WorkflowAuditor{logWriter: writer}

err := auditor.Close()

require.ErrorIs(t, err, closeErr)
assert.True(t, writer.closed)
})
}

func TestAuditor_CloseDoesNotCloseStdout(t *testing.T) {
t.Parallel()

auditor := &Auditor{logWriter: os.Stdout}

require.NoError(t, auditor.Close())
}

func TestWorkflowAuditor_LogWorkflowStarted(t *testing.T) {
t.Parallel()

Expand Down
21 changes: 21 additions & 0 deletions pkg/vmcp/core/core_vmcp.go
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,10 @@ type coreVMCP struct {
// by advertised tool name.
workflowDefs map[string]*composer.WorkflowDefinition

// workflowAuditor owns the optional workflow audit log writer and is closed
// with the core. Nil when workflow audit logging is disabled.
workflowAuditor *audit.WorkflowAuditor

// composerFactory builds a per-call composite-tool engine bound to a routing
// table, generalizing server.New's sessionComposerFactory (server.go:393).
composerFactory func(sessionRT *vmcp.RoutingTable, sessionTools []vmcp.Tool) composer.Composer
Expand Down Expand Up @@ -138,6 +142,14 @@ func New(cfg *Config) (VMCP, error) {
}
slog.Info("workflow audit logging enabled")
}
closeWorkflowAuditor := func() {
if workflowAuditor == nil {
return
}
if err := workflowAuditor.Close(); err != nil {
slog.Warn("failed to close workflow auditor", "error", err)
}
}

// The elicitation handler depends only on the domain-typed ElicitationRequester
// (#5436); no mcp-go types cross this boundary (vmcp anti-pattern #5).
Expand Down Expand Up @@ -168,6 +180,7 @@ func New(cfg *Config) (VMCP, error) {
instruments, err := newWorkflowInstruments(cfg.TelemetryProvider)
if err != nil {
stopStore()
closeWorkflowAuditor()
return nil, fmt.Errorf("failed to create workflow telemetry instruments: %w", err)
}

Expand Down Expand Up @@ -195,6 +208,7 @@ func New(cfg *Config) (VMCP, error) {
workflowDefs, err := validateWorkflowDefs(validationEngine, cfg.WorkflowDefs)
if err != nil {
stopStore()
closeWorkflowAuditor()
return nil, fmt.Errorf("workflow validation failed: %w", err)
}

Expand All @@ -206,6 +220,7 @@ func New(cfg *Config) (VMCP, error) {
healthMonitor, healthProvider, err := buildHealthMonitor(cfg)
if err != nil {
stopStore()
closeWorkflowAuditor()
return nil, err
}

Expand All @@ -217,6 +232,7 @@ func New(cfg *Config) (VMCP, error) {
healthMonitor: healthMonitor,
admission: admission,
workflowDefs: workflowDefs,
workflowAuditor: workflowAuditor,
composerFactory: composerFactory,
stopStore: stopStore,
}, nil
Expand Down Expand Up @@ -543,6 +559,11 @@ func (c *coreVMCP) InvalidateCapabilityCache() {
func (c *coreVMCP) Close() error {
c.closeOnce.Do(func() {
c.stopStore()
if c.workflowAuditor != nil {
if err := c.workflowAuditor.Close(); err != nil {
slog.Warn("failed to close workflow auditor", "error", err)
}
}
if c.healthMonitor != nil {
if err := c.healthMonitor.Stop(); err != nil {
slog.Warn("failed to stop health monitor", "error", err)
Expand Down