Skip to content
2 changes: 1 addition & 1 deletion server/cmd/gram/start.go
Original file line number Diff line number Diff line change
Expand Up @@ -1409,7 +1409,7 @@ func newStartCommand() *cli.Command {
roleManager := access.NewRoleManager(logger, db, roleClient, auditLogger)
access.Attach(mux, access.NewService(logger, tracerProvider, db, chDB, sessionManager, roleManager, authzEngine, auditLogger, emailService, siteURL, telemSvc))
agent.Attach(mux, agent.NewService(logger, tracerProvider, db, sessionManager, authzEngine, auditLogger, productFeatures, serverURL.String(), assetStorage, telemLogger, growthEmitter))
agentmanagement.Attach(mux, agentmanagement.NewService(logger, tracerProvider, db, sessionManager, authzEngine, auditLogger))
agentmanagement.Attach(mux, agentmanagement.NewService(logger, tracerProvider, db, sessionManager, authzEngine, auditLogger, featureFlags))
assistants.Attach(mux, assistantsSvc)
assistantmemories.Attach(mux, assistantmemories.NewService(
logger,
Expand Down
146 changes: 146 additions & 0 deletions server/internal/agentmanagement/gate_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,146 @@
package agentmanagement
Comment thread
cubic-dev-ai[bot] marked this conversation as resolved.

import (
"context"
"errors"
"testing"

"github.com/stretchr/testify/require"
"goa.design/goa/v3/security"

gen "github.com/speakeasy-api/gram/server/gen/agents"
"github.com/speakeasy-api/gram/server/internal/contextvalues"
"github.com/speakeasy-api/gram/server/internal/feature"
"github.com/speakeasy-api/gram/server/internal/oops"
"github.com/speakeasy-api/gram/server/internal/testenv"
)

type staticSessionAuthorizer struct {
authCtx *contextvalues.AuthContext
}

func (a staticSessionAuthorizer) AuthorizeWithPostAuthenticationCheck(
ctx context.Context,
_ string,
_ *security.APIKeyScheme,
check func(context.Context) error,
) (context.Context, error) {
ctx = contextvalues.SetAuthContext(ctx, a.authCtx)
return ctx, check(ctx)
}

type recordingAgentManagementFeatures struct {
evaluation feature.Evaluation
err error
flag feature.Flag
distinctID string
groups map[string]string
}

func (*recordingAgentManagementFeatures) IsFlagEnabled(context.Context, feature.Flag, string, map[string]string) (bool, error) {
return false, nil
}

func (*recordingAgentManagementFeatures) IsFlagEnabledLocal(context.Context, feature.Flag, string, map[string]string, map[string]string) (bool, error) {
return false, nil
}

func (*recordingAgentManagementFeatures) FlagPayload(context.Context, feature.Flag, string, map[string]string) ([]byte, error) {
return nil, nil
}

func (f *recordingAgentManagementFeatures) EvaluateFlag(_ context.Context, flag feature.Flag, distinctID string, groups map[string]string) (feature.Evaluation, error) {
f.flag = flag
f.distinctID = distinctID
f.groups = groups
return f.evaluation, f.err
}

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

backendFailure := errors.New("feature provider unavailable")
for _, test := range []struct {
name string
features feature.Provider
wantErr bool
wantLookup bool
}{
{name: "enabled", features: &recordingAgentManagementFeatures{evaluation: feature.EvaluationEnabled}, wantLookup: true},
{name: "disabled", features: &recordingAgentManagementFeatures{evaluation: feature.EvaluationDisabled}, wantErr: true, wantLookup: true},
{name: "indeterminate", features: &recordingAgentManagementFeatures{evaluation: feature.EvaluationIndeterminate}, wantErr: true, wantLookup: true},
{name: "provider error", features: &recordingAgentManagementFeatures{err: backendFailure}, wantErr: true, wantLookup: true},
{name: "missing provider", wantErr: true},
} {
t.Run(test.name, func(t *testing.T) {
t.Parallel()

service := &Service{logger: testenv.NewLogger(t), features: test.features}
ctx := contextvalues.SetAuthContext(t.Context(), &contextvalues.AuthContext{
ActiveOrganizationID: "organization",
OrganizationSlug: "organization-slug",
})

err := service.requireAgentManagementEnabled(ctx)
if test.wantErr {
requireOopsCode(t, err, oops.CodeNotFound)
} else {
require.NoError(t, err)
}

flags, ok := test.features.(*recordingAgentManagementFeatures)
if !test.wantLookup {
require.False(t, ok)
return
}
require.True(t, ok)
require.Equal(t, feature.FlagAgentManagement, flags.flag)
require.Equal(t, "organization", flags.distinctID)
require.Equal(t, feature.OrgProjectGroups("organization-slug", ""), flags.groups)
})
}
}

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

backendFailure := errors.New("feature provider unavailable")
for _, test := range []struct {
name string
features feature.Provider
}{
{name: "disabled", features: &recordingAgentManagementFeatures{evaluation: feature.EvaluationDisabled}},
{name: "indeterminate", features: &recordingAgentManagementFeatures{evaluation: feature.EvaluationIndeterminate}},
{name: "provider error", features: &recordingAgentManagementFeatures{err: backendFailure}},
{name: "missing provider"},
} {
t.Run(test.name, func(t *testing.T) {
t.Parallel()

authCtx := &contextvalues.AuthContext{
ActiveOrganizationID: "organization",
OrganizationSlug: "organization-slug",
}
service := &Service{
logger: testenv.NewLogger(t),
auth: staticSessionAuthorizer{authCtx: authCtx},
features: test.features,
}
endpoint := gen.NewCreateEndpoint(service, service.APIKeyAuth)

_, err := endpoint(t.Context(), &gen.CreatePayload{Name: "must not be created"})
requireOopsCode(t, err, oops.CodeNotFound)
})
}
}

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

features := &recordingAgentManagementFeatures{evaluation: feature.EvaluationEnabled}
service := &Service{logger: testenv.NewLogger(t), features: features}

requireOopsCode(t, service.requireAgentManagementEnabled(t.Context()), oops.CodeNotFound)
requireOopsCode(t, service.requireAgentManagementEnabled(contextvalues.SetAuthContext(t.Context(), &contextvalues.AuthContext{})), oops.CodeNotFound)
require.Empty(t, features.flag)
}
2 changes: 2 additions & 0 deletions server/internal/agentmanagement/ownership_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,7 @@ func TestTransferAtomicallyReplacesOwnerAndPreservesDirectPolicy(t *testing.T) {
require.Equal(t, "owner", actorID)
require.Equal(t, "owner", beforeOwner)
require.Equal(t, "replacement", afterOwner)
require.Equal(t, []string{"agent:create", "agent:policy_grant_create", "agent:transfer"}, agentWebhookOutboxActions(t, conn, "org-a"))
}

func TestExplicitReassignmentIsTheOnlyOwnershipOperationThatClearsLatch(t *testing.T) {
Expand Down Expand Up @@ -109,6 +110,7 @@ func TestExplicitReassignmentIsTheOnlyOwnershipOperationThatClearsLatch(t *testi
}
require.NoError(t, rows.Err())
require.Equal(t, []string{"agent:owner_loss", "agent:reassign"}, actions)
require.Equal(t, actions, agentWebhookOutboxActions(t, conn, "org-a"))
}

func TestOwnerChangeRejectsIneligibleAndCrossTenantTargets(t *testing.T) {
Expand Down
1 change: 1 addition & 0 deletions server/internal/agentmanagement/policy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,7 @@ func TestAgentPolicyCRUDIsExactAllowOnlyAndAudited(t *testing.T) {
}
require.NoError(t, rows.Err())
require.Equal(t, []string{"agent:policy_grant_create", "agent:policy_grant_update", "agent:policy_grant_delete"}, actions)
require.Equal(t, actions, agentWebhookOutboxActions(t, conn, "org-policy"))
}

func TestAgentPolicyRejectsUnsafeDenyAndMalformedGrantsAtomically(t *testing.T) {
Expand Down
44 changes: 42 additions & 2 deletions server/internal/agentmanagement/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,18 +28,29 @@ import (
"github.com/speakeasy-api/gram/server/internal/auth/sessions"
"github.com/speakeasy-api/gram/server/internal/authz"
"github.com/speakeasy-api/gram/server/internal/contextvalues"
"github.com/speakeasy-api/gram/server/internal/feature"
"github.com/speakeasy-api/gram/server/internal/middleware"
"github.com/speakeasy-api/gram/server/internal/oops"
"github.com/speakeasy-api/gram/server/internal/urn"
)

type sessionAuthorizer interface {
AuthorizeWithPostAuthenticationCheck(
context.Context,
string,
*security.APIKeyScheme,
func(context.Context) error,
) (context.Context, error)
}

type Service struct {
tracer trace.Tracer
logger *slog.Logger
db *pgxpool.Pool
auth *auth.Auth
auth sessionAuthorizer
authorizer *Authorizer
audit *audit.Logger
features feature.Provider
}

var _ gen.Service = (*Service)(nil)
Expand All @@ -52,6 +63,7 @@ func NewService(
sessionManager *sessions.Manager,
authzEngine *authz.Engine,
auditLogger *audit.Logger,
features feature.Provider,
) *Service {
logger = logger.With(attr.SlogComponent("agents"))
return &Service{
Expand All @@ -61,6 +73,7 @@ func NewService(
auth: auth.New(logger, db, sessionManager, authzEngine),
authorizer: NewAuthorizer(authzEngine),
audit: auditLogger,
features: features,
}
}

Expand All @@ -72,7 +85,34 @@ func Attach(mux goahttp.Muxer, service *Service) {
}

func (s *Service) APIKeyAuth(ctx context.Context, key string, schema *security.APIKeyScheme) (context.Context, error) {
return s.auth.Authorize(ctx, key, schema)
authorizedCtx, err := s.auth.AuthorizeWithPostAuthenticationCheck(ctx, key, schema, s.requireAgentManagementEnabled)
if err != nil {
return authorizedCtx, fmt.Errorf("authorize agent management session: %w", err)
}
return authorizedCtx, nil
}

func (s *Service) requireAgentManagementEnabled(ctx context.Context) error {
authCtx, ok := contextvalues.GetAuthContext(ctx)
if !ok || authCtx.ActiveOrganizationID == "" {
return oops.C(oops.CodeNotFound)
}

evaluation, err := feature.EvaluateFlag(
ctx,
s.features,
feature.FlagAgentManagement,
authCtx.ActiveOrganizationID,
feature.OrgProjectGroups(authCtx.OrganizationSlug, ""),
)
if err != nil {
s.logger.WarnContext(ctx, "failed to evaluate agent management rollout flag", attr.SlogError(err))
}
if evaluation != feature.EvaluationEnabled {
return oops.C(oops.CodeNotFound)
}

return nil
}

func (s *Service) Create(ctx context.Context, payload *gen.CreatePayload) (*gen.ManagedAgent, error) {
Expand Down
1 change: 1 addition & 0 deletions server/internal/agentmanagement/service_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,7 @@ func TestServiceLifecycleMutationsAreAuditedAtomically(t *testing.T) {
"agent:revoke",
"agent:delete",
}, actions)
require.Equal(t, actions, agentWebhookOutboxActions(t, conn, "org-a"))
}

func TestAuditFailureRollsBackLifecycleMutation(t *testing.T) {
Expand Down
35 changes: 35 additions & 0 deletions server/internal/agentmanagement/setup_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,19 +2,25 @@ package agentmanagement

import (
"context"
"encoding/json"
"log"
"log/slog"
"os"
"testing"

"github.com/jackc/pgx/v5/pgxpool"
"github.com/stretchr/testify/require"
"google.golang.org/protobuf/proto"

webhooksv1 "github.com/speakeasy-api/gram/infra/gen/gram/webhooks/v1"

"github.com/speakeasy-api/gram/server/internal/agents/repo"
"github.com/speakeasy-api/gram/server/internal/audit"
"github.com/speakeasy-api/gram/server/internal/conv"
orgrepo "github.com/speakeasy-api/gram/server/internal/organizations/repo"
"github.com/speakeasy-api/gram/server/internal/outbox/events"
"github.com/speakeasy-api/gram/server/internal/testenv"
"github.com/speakeasy-api/gram/server/internal/testenv/testrepo"
usersrepo "github.com/speakeasy-api/gram/server/internal/users/repo"
)

Expand Down Expand Up @@ -71,6 +77,35 @@ func createAgent(t *testing.T, conn *pgxpool.Pool, organizationID, ownerUserID,
return agent
}

func agentWebhookOutboxActions(t *testing.T, conn *pgxpool.Pool, organizationID string) []string {
t.Helper()

rows, err := testrepo.New(conn).ListPublishOutboxRows(t.Context())
require.NoError(t, err)

actions := make([]string, 0)
for _, row := range rows {
if row.OrganizationID != organizationID {
continue
}
var attributes map[string]string
require.NoError(t, json.Unmarshal(row.Attributes, &attributes))
if attributes["event_type"] != string(events.AgentV1.EventType()) {
continue
}

var event webhooksv1.Event
require.NoError(t, proto.Unmarshal(row.Message, &event))
var payload events.AuditLogCreatedPayloadV1
require.NoError(t, json.Unmarshal(event.GetPayload(), &payload))
require.Equal(t, organizationID, payload.OrganizationID)
require.Equal(t, "agent", payload.SubjectType)
actions = append(actions, payload.Action)
}

return actions
}

func newTestService(conn *pgxpool.Pool, engine authorizationEngine) *Service {
return &Service{
logger: slog.New(slog.NewTextHandler(os.Stderr, nil)),
Expand Down
26 changes: 26 additions & 0 deletions server/internal/auth/auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,24 @@ func New(logger *slog.Logger, db *pgxpool.Pool, sessions *sessions.Manager, auth
}

func (s *Auth) Authorize(ctx context.Context, key string, scheme *security.APIKeyScheme) (context.Context, error) {
return s.authorize(ctx, key, scheme, nil)
}

func (s *Auth) AuthorizeWithPostAuthenticationCheck(
ctx context.Context,
key string,
scheme *security.APIKeyScheme,
check func(context.Context) error,
) (context.Context, error) {
return s.authorize(ctx, key, scheme, check)
}

func (s *Auth) authorize(
ctx context.Context,
key string,
scheme *security.APIKeyScheme,
postAuthenticationCheck func(context.Context) error,
) (context.Context, error) {
if scheme == nil {
panic("Goa has not passed a schema") // TODO: figure something out here
}
Expand All @@ -64,6 +82,14 @@ func (s *Auth) Authorize(ctx context.Context, key string, scheme *security.APIKe
if err != nil {
return ctx, err
}

if postAuthenticationCheck != nil {
err = postAuthenticationCheck(ctx)
if err != nil {
return ctx, err
}
}

ctx, err = s.authz.PrepareContext(ctx)
if err != nil {
return ctx, oops.E(oops.CodeUnexpected, err, "load access grants").LogError(ctx, s.logger)
Expand Down
4 changes: 4 additions & 0 deletions server/internal/feature/flags.go
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,10 @@ const (
// FlagRiskEnforcementPubsub routes realtime gitleaks and Presidio scans over Pub/Sub.
FlagRiskEnforcementPubsub Flag = "risk-enforcement-pubsub"

// FlagAgentManagement gates the first-class agent management API. It is
// evaluated per organization and fails closed unless explicitly on.
FlagAgentManagement Flag = "agent-management"

// FlagDeviceLevelCoverage switches device-agent coverage from matching a
// device's assigned-user email against user-keyed heartbeats to matching
// its hardware serial against device-keyed ones, falling back to email
Expand Down
Loading
Loading