diff --git a/internal/guard/pipeline.go b/internal/guard/pipeline.go index 851b85b28..6a480b6d6 100644 --- a/internal/guard/pipeline.go +++ b/internal/guard/pipeline.go @@ -2,6 +2,7 @@ package guard import ( "context" + "errors" "fmt" "github.com/github/gh-aw-mcpg/internal/difc" @@ -64,6 +65,24 @@ func (e *PipelineAccessDenied) Error() string { return fmt.Sprintf("DIFC access denied: %s", e.EvalResult.Reason) } +// HandlePrePhaseError classifies a RunPipelinePrePhases error. +// +// When err represents a coarse-grained access denial, it returns the canonical +// *PipelineAccessDenied value together with the fully formatted violation error. +// For non-denial errors, it returns (nil, nil). +func HandlePrePhaseError(err error) (*PipelineAccessDenied, error) { + var denied *PipelineAccessDenied + if !errors.As(err, &denied) { + return nil, nil + } + return denied, difc.FormatViolationError( + denied.EvalResult, + denied.AgentLabels.Secrecy, + denied.AgentLabels.Integrity, + denied.Resource, + ) +} + // RunPipelinePrePhases executes phases 0–2 of the DIFC enforcement pipeline: // // - Phase 0: Get or create agent labels from the registry and store tool args in context. diff --git a/internal/guard/pipeline_test.go b/internal/guard/pipeline_test.go index 399aa01e9..3adfd40fd 100644 --- a/internal/guard/pipeline_test.go +++ b/internal/guard/pipeline_test.go @@ -3,6 +3,7 @@ package guard import ( "context" "errors" + "fmt" "testing" "github.com/stretchr/testify/assert" @@ -299,3 +300,31 @@ func TestPipelineAccessDenied_ErrorMessage(t *testing.T) { assert.Contains(t, denied.Error(), "DIFC access denied") assert.Contains(t, denied.Error(), "integrity tag missing") } + +func TestHandlePrePhaseError_ReturnsDetailedViolationForDeniedError(t *testing.T) { + agentLabels := difc.NewAgentLabels("test-agent") + agentLabels.AddSecrecyTags([]difc.Tag{"secret"}) + + resource := difc.NewLabeledResource("resource") + deniedErr := &PipelineAccessDenied{ + EvalResult: &difc.EvaluationResult{ + Decision: difc.AccessDeny, + SecrecyToAdd: []difc.Tag{"private"}, + Reason: "secrecy violation", + }, + Resource: resource, + AgentLabels: agentLabels, + } + + denied, detailedErr := HandlePrePhaseError(fmt.Errorf("wrapped: %w", deniedErr)) + assert.Same(t, deniedErr, denied) + assert.Error(t, detailedErr) + assert.Contains(t, detailedErr.Error(), "DIFC Violation:") + assert.Contains(t, detailedErr.Error(), "secrecy violation") +} + +func TestHandlePrePhaseError_IgnoresNonDeniedError(t *testing.T) { + denied, detailedErr := HandlePrePhaseError(errors.New("resource labeling failed")) + assert.Nil(t, denied) + assert.Nil(t, detailedErr) +} diff --git a/internal/proxy/handler.go b/internal/proxy/handler.go index e83b891de..495e5717f 100644 --- a/internal/proxy/handler.go +++ b/internal/proxy/handler.go @@ -205,7 +205,7 @@ func (h *proxyHandler) handleWithDIFC(w http.ResponseWriter, r *http.Request, pa } ctx, pre, err := guard.RunPipelinePrePhases(ctx, pipelineIn) if err != nil { - if denied, ok := err.(*guard.PipelineAccessDenied); ok { + if denied, _ := guard.HandlePrePhaseError(err); denied != nil { logHandler.Printf("[DIFC] Phase 2: BLOCKED %s %s — %s", r.Method, path, denied.EvalResult.Reason) deniedErr := fmt.Errorf("DIFC policy violation: %s", denied.EvalResult.Reason) tracing.RecordSpanError(difcSpan, deniedErr, "access denied: "+denied.EvalResult.Reason) diff --git a/internal/server/unified.go b/internal/server/unified.go index ebe069f09..b50c004e0 100644 --- a/internal/server/unified.go +++ b/internal/server/unified.go @@ -428,11 +428,9 @@ func (us *UnifiedServer) callBackendTool(ctx context.Context, serverID, toolName } ctx, pre, err := guard.RunPipelinePrePhases(ctx, pipelineIn) if err != nil { - if denied, ok := err.(*guard.PipelineAccessDenied); ok { + if denied, detailedErr := guard.HandlePrePhaseError(err); denied != nil { logger.LogWarn("difc", "Access DENIED for agent %s to %s: %s", agentID, denied.Resource.Description, denied.EvalResult.Reason) - detailedErr := difc.FormatViolationError(denied.EvalResult, - denied.AgentLabels.Secrecy, denied.AgentLabels.Integrity, denied.Resource) tracing.RecordSpanError(toolSpan, detailedErr, "access denied: "+denied.EvalResult.Reason) httpStatusCode = 403 return mcp.NewErrorCallToolResult(detailedErr)