diff --git a/server/internal/mcp/audit.go b/server/internal/mcp/audit.go new file mode 100644 index 0000000..9f75285 --- /dev/null +++ b/server/internal/mcp/audit.go @@ -0,0 +1,95 @@ +package mcp + +import ( + "errors" + "fmt" + "sort" + "strings" + + "gitea.hostxtra.co.uk/mrhid6/vantage/server/internal/services" +) + +// FanOutLimit is how many servers a write tool may touch before it demands +// explicit confirmation. Cheap insurance against a mis-parsed selector reaching +// the whole fleet on one badly phrased instruction. +const FanOutLimit = 25 + +// ErrConfirmRequired is returned to the model as a tool error it can act on: +// it says what would have happened and how to proceed deliberately. +var ErrConfirmRequired = errors.New("confirmation required") + +// CheckFanOut refuses a write that would touch more servers than FanOutLimit +// unless the call passed confirm:true. +func CheckFanOut(count int, args map[string]any) error { + if count <= FanOutLimit { + return nil + } + if confirm, ok := args["confirm"].(bool); ok && confirm { + return nil + } + return fmt.Errorf("%w: this would affect %d servers, above the limit of %d; "+ + "call again with confirm:true if that is intended", + ErrConfirmRequired, count, FanOutLimit) +} + +// SummariseArgs renders an argument object as a short, deterministic, +// bounded string for the audit log. Values are described rather than +// reproduced: an argument may carry arbitrary text a model generated. +func SummariseArgs(args map[string]any) string { + if len(args) == 0 { + return "no arguments" + } + + keys := make([]string, 0, len(args)) + for k := range args { + keys = append(keys, k) + } + sort.Strings(keys) + + parts := make([]string, 0, len(keys)) + for _, k := range keys { + parts = append(parts, k+"="+summariseValue(args[k])) + } + + out := strings.Join(parts, " ") + if len(out) > 200 { + out = out[:197] + "..." + } + return out +} + +func summariseValue(v any) string { + switch t := v.(type) { + case string: + if len(t) > 40 { + return fmt.Sprintf("<%d chars>", len(t)) + } + return t + case bool, float64, int: + return fmt.Sprint(t) + case []any: + return fmt.Sprintf("<%d items>", len(t)) + case map[string]any: + return fmt.Sprintf("<%d fields>", len(t)) + default: + return "" + } +} + +// LogCall records a successful tool call. Reads are recorded as well as writes: +// the point of an agent-facing surface is being able to reconstruct afterwards +// what the agent looked at, not only what it changed. +func LogCall(c Caller, t Tool, args map[string]any, servers int) { + detail := fmt.Sprintf("tool %s (%s)", t.Name, SummariseArgs(args)) + if servers > 0 { + detail += fmt.Sprintf(", %d server(s) affected", servers) + } + services.LogEvent(c.InstanceID, "mcp.tool_call", c.TokenName, "", "", detail) +} + +// LogDenied records a refusal and which gate refused, which is what turns "the +// agent said it could not" into a diagnosable event. +func LogDenied(c Caller, toolName, gate string) { + services.LogEvent(c.InstanceID, "mcp.tool_denied", c.TokenName, "", "", + fmt.Sprintf("tool %s refused by %s", toolName, gate)) +} diff --git a/server/internal/mcp/audit_test.go b/server/internal/mcp/audit_test.go new file mode 100644 index 0000000..22d0949 --- /dev/null +++ b/server/internal/mcp/audit_test.go @@ -0,0 +1,45 @@ +package mcp + +import ( + "strings" + "testing" +) + +// Arguments can carry arbitrary model output and the audit log is read by +// humans in a UI, so they are summarised rather than dumped. +func TestSummariseArgsIsBoundedAndOrdered(t *testing.T) { + got := SummariseArgs(map[string]any{ + "workflow_id": "wf-1", + "note": strings.Repeat("x", 500), + }) + + if len(got) > 200 { + t.Errorf("summary is %d chars, want at most 200", len(got)) + } + if !strings.Contains(got, "workflow_id") { + t.Errorf("summary %q omits an argument name", got) + } + // Deterministic ordering, or two identical calls produce different audit + // rows and nothing can be compared. + if SummariseArgs(map[string]any{"b": 1, "a": 2}) != SummariseArgs(map[string]any{"a": 2, "b": 1}) { + t.Error("SummariseArgs is not deterministic") + } +} + +func TestCheckFanOutRequiresConfirmation(t *testing.T) { + if err := CheckFanOut(5, nil); err != nil { + t.Errorf("CheckFanOut(5) = %v, want nil", err) + } + + err := CheckFanOut(200, nil) + if err == nil { + t.Fatal("CheckFanOut(200) = nil, want a refusal") + } + if !strings.Contains(err.Error(), "200") { + t.Errorf("refusal %q does not say how many servers", err) + } + + if err := CheckFanOut(200, map[string]any{"confirm": true}); err != nil { + t.Errorf("CheckFanOut(200, confirm) = %v, want nil", err) + } +}