Skip to content
36 changes: 21 additions & 15 deletions cmd/ateapi/internal/actoridentity/actoridentity.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ import (
"time"

"github.com/agent-substrate/substrate/cmd/ateapi/internal/actoridjwt"
"github.com/agent-substrate/substrate/cmd/ateapi/internal/controlapi"
"github.com/agent-substrate/substrate/cmd/ateapi/internal/store"
"github.com/agent-substrate/substrate/cmd/ateapi/internal/workercache"
"github.com/agent-substrate/substrate/internal/localca"
Expand All @@ -40,6 +41,7 @@ import (
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/peer"
"google.golang.org/grpc/status"
"k8s.io/apimachinery/pkg/api/operation"
"k8s.io/apimachinery/pkg/util/validation/field"
)

Expand Down Expand Up @@ -94,6 +96,10 @@ func (s *Server) MintJWT(ctx context.Context, req *ateapipb.MintJWTRequest) (*at
return nil, status.Errorf(codes.PermissionDenied, "caller is not permitted to mint actor JWTs")
}

if errs := validateMintJWTRequest(ctx, req); len(errs) > 0 {
return nil, status.Error(codes.InvalidArgument, errs.ToAggregate().Error())
}

// TODO: Cross-check the verified caller and requested actor against the actor database.

// TODO: Cache signing keys in memory, so we don't read from disk every time.
Expand All @@ -106,10 +112,6 @@ func (s *Server) MintJWT(ctx context.Context, req *ateapipb.MintJWTRequest) (*at
if err != nil {
return nil, fmt.Errorf("while unmarshaling signing pool: %w", err)
}
// We only issue tokens with audience bindings.
if len(req.GetAudience()) == 0 {
return nil, fmt.Errorf("at least one audience must be requested")
}

actorClaims := &actoridjwt.Claims{
// TODO: This is currently API but it has to be a globally unique, oidc-compliant and accsible DNS name
Expand Down Expand Up @@ -150,16 +152,14 @@ func (s *Server) MintCert(ctx context.Context, req *ateapipb.MintCertRequest) (*
if err != nil {
return nil, err
}
if errs := validateMintCertRequest(ctx, req); len(errs) > 0 {
return nil, status.Error(codes.InvalidArgument, errs.ToAggregate().Error())
}
// Validation bounds purpose to the enum's range; which purposes this
// server actually supports is a policy decision that stays here.
if req.GetPurpose() != ateapipb.ActorCertificatePurpose_ACTOR_CERTIFICATE_PURPOSE_ATUNNEL {
return nil, status.Error(codes.InvalidArgument, "unsupported actor certificate purpose")
}

if err := validateWorkerRef(req.GetWorker()); err != nil {
return nil, status.Errorf(codes.InvalidArgument, "invalid worker: %v", err)
}
if req.GetExpectedActorUid() == "" {
return nil, status.Error(codes.InvalidArgument, "expected_actor_uid is required")
}
actor, actorRef, err := s.authorizeActor(ctx, caller, req)
if err != nil {
return nil, err
Expand Down Expand Up @@ -286,10 +286,16 @@ func authenticateAtelet(ctx context.Context) (*ateletCaller, error) {
return &ateletCaller{podName: identity.PodName, nodeName: identity.NodeName}, nil
}

// validateWorkerRef checks the reference to the Worker the certificate is
// minted for. Workers are global-scoped, so the reference carries no atespace.
func validateWorkerRef(worker *ateapipb.ObjectRef) error {
return resources.ValidateGlobalObjectRef(worker, field.NewPath("worker")).ToAggregate()
func validateMintJWTRequest(ctx context.Context, req *ateapipb.MintJWTRequest) field.ErrorList {
// Call the generated validation.
op := operation.Operation{Type: operation.Create}
return controlapi.Validate_MintJWTRequest(ctx, op, nil, req, nil)
}

func validateMintCertRequest(ctx context.Context, req *ateapipb.MintCertRequest) field.ErrorList {
// Call the generated validation.
op := operation.Operation{Type: operation.Create}
return controlapi.Validate_MintCertRequest(ctx, op, nil, req, nil)
}

// authorizeActor resolves the actor from the authenticated worker and verifies
Expand Down
161 changes: 160 additions & 1 deletion cmd/ateapi/internal/actoridentity/actoridentity_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,9 +21,11 @@ import (
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"fmt"
"math/big"
"net/url"
"path"
"strings"
"testing"
"time"

Expand All @@ -39,8 +41,14 @@ import (
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/peer"
"google.golang.org/grpc/status"
"k8s.io/apimachinery/pkg/util/validation/field"
)

func assertValidateErr(t *testing.T, got field.ErrorList, want field.ErrorList) {
t.Helper()
field.ErrorMatcher{}.ByType().ByField().ByOrigin().Test(t, want, got)
}

const (
testAtespace = "team-alpha"
testActorName = "counter-1"
Expand Down Expand Up @@ -777,7 +785,7 @@ func TestMintCertActorUID(t *testing.T) {
wantCode codes.Code
}{
"Matching": {requestUID: func(actorUID string) string { return actorUID }, wantCode: codes.OK},
"Stale": {requestUID: func(string) string { return "uid-of-a-previous-incarnation" }, wantCode: codes.FailedPrecondition},
"Stale": {requestUID: func(string) string { return "9d1f7b06-3c58-4a2e-8b40-5f7c1e9a2d63" }, wantCode: codes.FailedPrecondition},
} {
t.Run(name, func(t *testing.T) {
leaf, actorUID, err := mintCertFor(t, func(actorUID string) *ateapipb.MintCertRequest {
Expand Down Expand Up @@ -924,3 +932,154 @@ func TestMintCertAuthorizesBeforeSigning(t *testing.T) {
t.Errorf("MintCert() code = %v (err = %v), want %v", got, err, codes.PermissionDenied)
}
}

func TestValidateMintJWTRequest(t *testing.T) {
// This test verifies validation of user input for minting a JWT.
validReq := func(mods ...func(req *ateapipb.MintJWTRequest)) *ateapipb.MintJWTRequest {
req := &ateapipb.MintJWTRequest{
Audience: []string{"aud1"},
Atespace: "as1",
ActorName: "actor1",
ActorUid: "01234567-89ab-cdef-0123-456789abcdef",
}
for _, m := range mods {
m(req)
}
return req
}

tests := []struct {
name string
req *ateapipb.MintJWTRequest
want field.ErrorList
}{{
"valid",
validReq(),
nil,
}, {
"missing audience",
validReq(func(r *ateapipb.MintJWTRequest) { r.Audience = nil }),
field.ErrorList{field.Required(field.NewPath("audience"), "")},
}, {
"too many audiences",
validReq(func(r *ateapipb.MintJWTRequest) {
r.Audience = make([]string, 17)
for i := range r.Audience {
r.Audience[i] = fmt.Sprintf("https://svc-%d.example.com", i)
}
}),
field.ErrorList{field.TooMany(field.NewPath("audience"), 17, 16).WithOrigin("maxItems")},
}, {
"duplicate audience entry",
validReq(func(r *ateapipb.MintJWTRequest) {
r.Audience = []string{"https://a.example.com", "https://a.example.com"}
}),
field.ErrorList{field.Duplicate(field.NewPath("audience").Index(1), nil)},
}, {
"audience entry too long",
validReq(func(r *ateapipb.MintJWTRequest) { r.Audience = []string{strings.Repeat("a", 513)} }),
field.ErrorList{field.TooLong(field.NewPath("audience").Index(0), nil, 512).WithOrigin("maxLength")},
}, {
"missing atespace",
validReq(func(r *ateapipb.MintJWTRequest) { r.Atespace = "" }),
field.ErrorList{field.Required(field.NewPath("atespace"), "")},
}, {
"invalid atespace",
validReq(func(r *ateapipb.MintJWTRequest) { r.Atespace = "AS1" }),
field.ErrorList{field.Invalid(field.NewPath("atespace"), nil, "").WithOrigin("format=k8s-short-name")},
}, {
"missing actor_name",
validReq(func(r *ateapipb.MintJWTRequest) { r.ActorName = "" }),
field.ErrorList{field.Required(field.NewPath("actor_name"), "")},
}, {
"invalid actor_name",
validReq(func(r *ateapipb.MintJWTRequest) { r.ActorName = "invalid value" }),
field.ErrorList{field.Invalid(field.NewPath("actor_name"), nil, "").WithOrigin("format=k8s-short-name")},
}, {
"unspecified actor_uid",
validReq(func(r *ateapipb.MintJWTRequest) { r.ActorUid = "" }),
nil,
}, {
"invalid actor_uid",
validReq(func(r *ateapipb.MintJWTRequest) { r.ActorUid = "not a uid" }),
field.ErrorList{field.Invalid(field.NewPath("actor_uid"), nil, "").WithOrigin("format=k8s-uuid")},
}}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assertValidateErr(t, validateMintJWTRequest(context.Background(), tt.req), tt.want)
})
}
}

func TestValidateMintCertRequest(t *testing.T) {
// This test verifies validation of user input for minting a certificate.
validReq := func(mods ...func(req *ateapipb.MintCertRequest)) *ateapipb.MintCertRequest {
req := &ateapipb.MintCertRequest{
Worker: &ateapipb.ObjectRef{Name: "worker1"},
CertificateSigningRequest: []byte{0x01},
ExpectedActorUid: "01234567-89ab-cdef-0123-456789abcdef",
Purpose: ateapipb.ActorCertificatePurpose_ACTOR_CERTIFICATE_PURPOSE_ATUNNEL,
}
for _, m := range mods {
m(req)
}
return req
}

tests := []struct {
name string
req *ateapipb.MintCertRequest
want field.ErrorList
}{{
"valid",
validReq(),
nil,
}, {
"oversized certificate_signing_request",
validReq(func(r *ateapipb.MintCertRequest) { r.CertificateSigningRequest = make([]byte, 16385) }),
field.ErrorList{field.TooLong(field.NewPath("certificate_signing_request"), nil, 16384)},
}, {
"missing worker",
validReq(func(r *ateapipb.MintCertRequest) { r.Worker = nil }),
field.ErrorList{field.Required(field.NewPath("worker"), "")},
}, {
"worker.atespace must be empty",
validReq(func(r *ateapipb.MintCertRequest) { r.Worker.Atespace = "as1" }),
field.ErrorList{field.Forbidden(field.NewPath("worker", "atespace"), "")},
}, {
"missing worker.name",
validReq(func(r *ateapipb.MintCertRequest) { r.Worker.Name = "" }),
field.ErrorList{field.Required(field.NewPath("worker", "name"), "")},
}, {
"invalid worker.name",
validReq(func(r *ateapipb.MintCertRequest) { r.Worker.Name = "invalid value" }),
field.ErrorList{field.Invalid(field.NewPath("worker", "name"), nil, "").WithOrigin("format=k8s-short-name")},
}, {
"missing certificate_signing_request",
validReq(func(r *ateapipb.MintCertRequest) { r.CertificateSigningRequest = nil }),
field.ErrorList{field.Required(field.NewPath("certificate_signing_request"), "")},
}, {
"missing expected_actor_uid",
validReq(func(r *ateapipb.MintCertRequest) { r.ExpectedActorUid = "" }),
field.ErrorList{field.Required(field.NewPath("expected_actor_uid"), "")},
}, {
"invalid expected_actor_uid",
validReq(func(r *ateapipb.MintCertRequest) { r.ExpectedActorUid = "not a uid" }),
field.ErrorList{field.Invalid(field.NewPath("expected_actor_uid"), nil, "").WithOrigin("format=k8s-uuid")},
}, {
"unspecified purpose",
validReq(func(r *ateapipb.MintCertRequest) {
r.Purpose = ateapipb.ActorCertificatePurpose_ACTOR_CERTIFICATE_PURPOSE_UNSPECIFIED
}),
field.ErrorList{field.Required(field.NewPath("purpose"), "")},
}, {
"out-of-range purpose",
validReq(func(r *ateapipb.MintCertRequest) { r.Purpose = ateapipb.ActorCertificatePurpose(99) }),
field.ErrorList{field.Invalid(field.NewPath("purpose"), nil, "").WithOrigin("maximum")},
}}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assertValidateErr(t, validateMintCertRequest(context.Background(), tt.req), tt.want)
})
}
}
11 changes: 11 additions & 0 deletions cmd/ateapi/internal/controlapi/validate.go
Original file line number Diff line number Diff line change
Expand Up @@ -97,3 +97,14 @@ func ValidateCustom_UpdateActorRequest_Actor(ctx context.Context, op operation.O
func ValidateCustom_WorkerAssignment_WorkerPodIp(_ context.Context, _ operation.Operation, fldPath *field.Path, value, _ *string) field.ErrorList {
return validation.IsValidIP(fldPath, *value)
}

// maxCSRBytes bounds MintCertRequest's CSR. Real CSRs are a few KB; this is
// a guardrail, applied here because maxLength does not support bytes fields.
const maxCSRBytes = 16384

func ValidateCustom_MintCertRequest_CertificateSigningRequest(_ context.Context, _ operation.Operation, fldPath *field.Path, value, _ []byte) field.ErrorList {
if len(value) > maxCSRBytes {
return field.ErrorList{field.TooLong(fldPath, nil, maxCSRBytes)}
}
return nil
}
Loading
Loading