diff --git a/pkg/audit/auditor.go b/pkg/audit/auditor.go index 9a6d288668..b86c636bb1 100644 --- a/pkg/audit/auditor.go +++ b/pkg/audit/auditor.go @@ -88,7 +88,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 diff --git a/pkg/audit/auditor_test.go b/pkg/audit/auditor_test.go index 31ca2166c1..432930c816 100644 --- a/pkg/audit/auditor_test.go +++ b/pkg/audit/auditor_test.go @@ -11,6 +11,7 @@ import ( "io" "net/http" "net/http/httptest" + "os" "strings" "testing" "time" @@ -34,6 +35,16 @@ func TestNewAuditor(t *testing.T) { assert.Equal(t, config, auditor.config) } +func TestAuditor_CloseDoesNotCloseStdout(t *testing.T) { + t.Parallel() + + auditor := &Auditor{logWriter: os.Stdout} + + require.NoError(t, auditor.Close()) + _, err := os.Stdout.Write(nil) + require.NoError(t, err, "Close() must not close os.Stdout") +} + func TestAuditorMiddlewareDisabled(t *testing.T) { t.Parallel() config := &Config{} diff --git a/pkg/audit/workflow_auditor.go b/pkg/audit/workflow_auditor.go index e87a3ce94c..93769adc76 100644 --- a/pkg/audit/workflow_auditor.go +++ b/pkg/audit/workflow_auditor.go @@ -8,6 +8,7 @@ import ( "context" "encoding/json" "fmt" + "io" "log/slog" "time" @@ -21,6 +22,7 @@ type WorkflowAuditor struct { auditLogger *slog.Logger config *Config component string + logWriter io.Writer } // NewWorkflowAuditor creates a new workflow auditor. @@ -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, diff --git a/pkg/audit/workflow_auditor_test.go b/pkg/audit/workflow_auditor_test.go index 4b8f49dd2c..3218e9cf00 100644 --- a/pkg/audit/workflow_auditor_test.go +++ b/pkg/audit/workflow_auditor_test.go @@ -24,11 +24,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 "" @@ -53,6 +67,7 @@ func createTestAuditor(t *testing.T, config *Config) (*WorkflowAuditor, *testLog auditLogger: NewAuditLogger(writer), config: config, component: "vmcp-composer", + logWriter: writer, } return auditor, writer @@ -123,6 +138,48 @@ 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()) + _, err = os.Stdout.Write(nil) + require.NoError(t, err, "Close() must not close os.Stdout") + 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 TestWorkflowAuditor_LogWorkflowStarted(t *testing.T) { t.Parallel() diff --git a/pkg/vmcp/core/core_vmcp.go b/pkg/vmcp/core/core_vmcp.go index ec85b01fb5..494e2f3e9b 100644 --- a/pkg/vmcp/core/core_vmcp.go +++ b/pkg/vmcp/core/core_vmcp.go @@ -5,6 +5,7 @@ package core import ( "context" + "errors" "fmt" "log/slog" "sync" @@ -71,6 +72,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 @@ -81,6 +86,7 @@ type coreVMCP struct { stopStore func() closeOnce sync.Once + closeErr error } var _ VMCP = (*coreVMCP)(nil) @@ -138,6 +144,16 @@ func New(cfg *Config) (VMCP, error) { } slog.Info("workflow audit logging enabled") } + closeWorkflowAuditor := func() error { + if workflowAuditor == nil { + return nil + } + if err := workflowAuditor.Close(); err != nil { + slog.Warn("failed to close workflow auditor", "error", err) + return err + } + return nil + } // The elicitation handler depends only on the domain-typed ElicitationRequester // (#5436); no mcp-go types cross this boundary (vmcp anti-pattern #5). @@ -168,7 +184,11 @@ func New(cfg *Config) (VMCP, error) { instruments, err := newWorkflowInstruments(cfg.TelemetryProvider) if err != nil { stopStore() - return nil, fmt.Errorf("failed to create workflow telemetry instruments: %w", err) + cleanupErr := closeWorkflowAuditor() + return nil, errors.Join( + fmt.Errorf("failed to create workflow telemetry instruments: %w", err), + cleanupErr, + ) } // composerFactory builds a composite-tool engine bound to a specific routing @@ -195,7 +215,8 @@ func New(cfg *Config) (VMCP, error) { workflowDefs, err := validateWorkflowDefs(validationEngine, cfg.WorkflowDefs) if err != nil { stopStore() - return nil, fmt.Errorf("workflow validation failed: %w", err) + cleanupErr := closeWorkflowAuditor() + return nil, errors.Join(fmt.Errorf("workflow validation failed: %w", err), cleanupErr) } // Build and start the backend health monitor (#5443 reversal: the core owns its @@ -206,7 +227,8 @@ func New(cfg *Config) (VMCP, error) { healthMonitor, healthProvider, err := buildHealthMonitor(cfg) if err != nil { stopStore() - return nil, err + cleanupErr := closeWorkflowAuditor() + return nil, errors.Join(err, cleanupErr) } return &coreVMCP{ @@ -217,6 +239,7 @@ func New(cfg *Config) (VMCP, error) { healthMonitor: healthMonitor, admission: admission, workflowDefs: workflowDefs, + workflowAuditor: workflowAuditor, composerFactory: composerFactory, stopStore: stopStore, }, nil @@ -537,19 +560,32 @@ func (c *coreVMCP) InvalidateCapabilityCache() { invalidator.InvalidateAll() } -// Close stops the workflow state store's cleanup goroutine. It is idempotent: -// the underlying Stop closes a channel that cannot be closed twice, so the work -// is guarded by sync.Once and subsequent calls return nil. +// Close stops the workflow state store's cleanup goroutine, closes the optional +// workflow auditor, and stops the health monitor. It returns cleanup errors from +// the first call, and is idempotent: the underlying cleanup work is guarded by +// sync.Once and subsequent calls return nil. func (c *coreVMCP) Close() error { + ran := false c.closeOnce.Do(func() { + ran = true c.stopStore() + if c.workflowAuditor != nil { + if err := c.workflowAuditor.Close(); err != nil { + slog.Warn("failed to close workflow auditor", "error", err) + c.closeErr = errors.Join(c.closeErr, err) + } + } if c.healthMonitor != nil { if err := c.healthMonitor.Stop(); err != nil { slog.Warn("failed to stop health monitor", "error", err) + c.closeErr = errors.Join(c.closeErr, err) } } }) - return nil + if !ran { + return nil + } + return c.closeErr } // aggregatedView health-filters the backend registry and aggregates capabilities diff --git a/pkg/vmcp/core/core_vmcp_test.go b/pkg/vmcp/core/core_vmcp_test.go index 5506635fba..0d0ea11121 100644 --- a/pkg/vmcp/core/core_vmcp_test.go +++ b/pkg/vmcp/core/core_vmcp_test.go @@ -8,6 +8,9 @@ import ( "context" "errors" "log/slog" + "os" + "path/filepath" + "runtime" "testing" "time" @@ -15,6 +18,7 @@ import ( "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" + "github.com/stacklok/toolhive/pkg/audit" "github.com/stacklok/toolhive/pkg/auth" "github.com/stacklok/toolhive/pkg/vmcp" "github.com/stacklok/toolhive/pkg/vmcp/aggregator" @@ -61,6 +65,31 @@ func baseConfig(t *testing.T) (*Config, *coreMocks) { return cfg, m } +func requireNoOpenFDForPath(t *testing.T, path string) { + t.Helper() + + if runtime.GOOS != "linux" { + t.Skip("/proc/self/fd is Linux-only; fd-release is covered on Linux CI") + } + + entries, err := os.ReadDir("/proc/self/fd") + require.NoError(t, err) + for _, entry := range entries { + target, err := os.Readlink(filepath.Join("/proc/self/fd", entry.Name())) + if err != nil { + continue + } + require.NotEqual(t, path, target, "audit log file descriptor must be closed") + } +} + +func closeErrorPathAuditConfig(t *testing.T) (*audit.Config, string) { + t.Helper() + + path := filepath.Join(t.TempDir(), "workflow-audit.log") + return &audit.Config{LogFile: path}, path +} + // testBackendID is the single backend ID used across these tests. const testBackendID = "be1" @@ -115,6 +144,87 @@ func TestNew_NilConfig(t *testing.T) { assert.Nil(t, c) } +func TestNew_CloseReleasesWorkflowAuditLogFile(t *testing.T) { + t.Parallel() + + t.Run("closes retained audit log file", func(t *testing.T) { + t.Parallel() + + cfg, _ := baseConfig(t) + cfg.AuditConfig, _ = closeErrorPathAuditConfig(t) + + c, err := New(cfg) + require.NoError(t, err) + require.NoError(t, c.Close()) + + require.ErrorIs(t, c.(*coreVMCP).workflowAuditor.Close(), os.ErrClosed) + }) + + t.Run("returns first close error and keeps later Close idempotent", func(t *testing.T) { + t.Parallel() + + cfg, _ := baseConfig(t) + cfg.AuditConfig, _ = closeErrorPathAuditConfig(t) + + c, err := New(cfg) + require.NoError(t, err) + core := c.(*coreVMCP) + require.NoError(t, core.workflowAuditor.Close()) + + err = c.Close() + require.Error(t, err) + assert.ErrorIs(t, err, os.ErrClosed) + require.NoError(t, c.Close()) + }) +} + +func TestNew_ErrorPathsCloseWorkflowAuditLogFile(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + mutate func(*testing.T, *Config, *coreMocks) + }{ + { + name: "workflow validation error", + mutate: func(_ *testing.T, cfg *Config, _ *coreMocks) { + cfg.WorkflowDefs = map[string]*composer.WorkflowDefinition{ + "wf": { + Name: "wf", + Steps: []composer.WorkflowStep{ + {ID: "s1", Type: composer.StepTypeTool, Tool: "be1.tool", DependsOn: []string{"s2"}}, + {ID: "s2", Type: composer.StepTypeTool, Tool: "be1.tool", DependsOn: []string{"s1"}}, + }, + }, + } + }, + }, + { + name: "health monitor creation error", + mutate: func(_ *testing.T, cfg *Config, mocks *coreMocks) { + mocks.reg.EXPECT().List(gomock.Any()).Return(nil) + cfg.HealthMonitorConfig = &health.MonitorConfig{} + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + cfg, mocks := baseConfig(t) + var auditLogPath string + cfg.AuditConfig, auditLogPath = closeErrorPathAuditConfig(t) + tt.mutate(t, cfg, mocks) + + c, err := New(cfg) + + require.Error(t, err) + assert.Nil(t, c) + requireNoOpenFDForPath(t, auditLogPath) + }) + } +} + func TestNew_ValidatesWorkflows(t *testing.T) { t.Parallel()