Skip to content
527 changes: 348 additions & 179 deletions packages/api/internal/api/api.gen.go

Large diffs are not rendered by default.

74 changes: 72 additions & 2 deletions packages/api/internal/api/spec_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,8 @@ func TestSpecSecuritySchemeHeaderNames(t *testing.T) {
{"Supabase1TokenAuth", auth.HeaderSupabaseToken},
{"Supabase2TeamAuth", auth.HeaderSupabaseTeam},
{"AuthProviderTeamAuth", auth.HeaderTeamID},
{"AdminTokenAuth", auth.HeaderAdminToken},
{"AdminApiKeyAuth", auth.HeaderAdminToken},
{"AdminTeamAuth", auth.HeaderTeamID},
}

for _, tc := range cases {
Expand Down Expand Up @@ -77,7 +78,8 @@ func TestAuthProviderTeamAuthHeaderRoutes(t *testing.T) {
"Supabase2TeamAuth": auth.HeaderSupabaseTeam,
"AuthProviderBearerAuth": auth.HeaderAuthorization,
"AuthProviderTeamAuth": auth.HeaderTeamID,
"AdminTokenAuth": auth.HeaderAdminToken,
"AdminApiKeyAuth": auth.HeaderAdminToken,
"AdminTeamAuth": auth.HeaderTeamID,
}

authFn := func(_ context.Context, input *openapi3filter.AuthenticationInput) error {
Expand Down Expand Up @@ -145,3 +147,71 @@ func TestAuthProviderTeamAuthHeaderRoutes(t *testing.T) {
require.Equal(t, "AuthProviderTeamAuth", gotSchemeName)
require.Equal(t, wantToken, gotToken)
}

// TestAdminTeamAuthSchemeOrder verifies that the admin security
// schemes validate the admin token before the team header. kin-openapi sorts
// schemes by name inside one security requirement.
func TestAdminTeamAuthSchemeOrder(t *testing.T) {
t.Parallel()

swagger, err := GetSpec()
require.NoError(t, err)
swagger.Servers = nil

var adminSchemeOrder []string

schemeHeaders := map[string]string{
"ApiKeyAuth": auth.HeaderAPIKey,
"AccessTokenAuth": auth.HeaderAuthorization,
"Supabase1TokenAuth": auth.HeaderSupabaseToken,
"Supabase2TeamAuth": auth.HeaderSupabaseTeam,
"AuthProviderBearerAuth": auth.HeaderAuthorization,
"AuthProviderTeamAuth": auth.HeaderTeamID,
"AdminApiKeyAuth": auth.HeaderAdminToken,
"AdminTeamAuth": auth.HeaderTeamID,
}

authFn := func(_ context.Context, input *openapi3filter.AuthenticationInput) error {
header, ok := schemeHeaders[input.SecuritySchemeName]
if !ok {
return http.ErrNoCookie
}

if value := input.RequestValidationInput.Request.Header.Get(header); value == "" {
return http.ErrNoCookie
}

switch input.SecuritySchemeName {
case "AdminApiKeyAuth", "AdminTeamAuth":
adminSchemeOrder = append(adminSchemeOrder, input.SecuritySchemeName)
}

return nil
}

r := gin.New()
r.Use(middleware.OapiRequestValidatorWithOptions(swagger, &middleware.Options{
Options: openapi3filter.Options{
AuthenticationFunc: authFn,
MultiError: true,
},
SilenceServersWarning: true,
}))

r.NoRoute(func(c *gin.Context) {
c.Status(http.StatusOK)
})
r.Any("/*any", func(c *gin.Context) {
c.Status(http.StatusOK)
})

req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/api-keys", nil)
req.Header.Set(auth.HeaderAdminToken, "admin-token")
req.Header.Set(auth.HeaderTeamID, "team-id")

rr := httptest.NewRecorder()
r.ServeHTTP(rr, req)

require.Equal(t, http.StatusOK, rr.Code, "request with admin token and team header should pass auth (body: %s)", rr.Body.String())
require.Equal(t, []string{"AdminApiKeyAuth", "AdminTeamAuth"}, adminSchemeOrder)
}
40 changes: 39 additions & 1 deletion packages/api/internal/cfg/model.go
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,41 @@ type Config struct {
DomainName string `env:"DOMAIN_NAME" envDefault:""`
}

type FailureCondition string

const (
FailureConditionInvalidServiceDiscoveryProvider FailureCondition = "invalid_service_discovery_provider"
)

type FailureError struct {
Condition FailureCondition
err error
}

func (e *FailureError) Error() string {
return e.err.Error()
}

func (e *FailureError) Unwrap() error {
return e.err
}

func ParseFailureCondition(err error) (FailureCondition, bool) {
var failureErr *FailureError
if !errors.As(err, &failureErr) {
return "", false
}

return failureErr.Condition, true
}

func newFailureError(condition FailureCondition, message string) error {
return &FailureError{
Condition: condition,
err: errors.New(message),
}
}

type JWTSigningKey any

type VolumesTokenConfig struct {
Expand Down Expand Up @@ -157,7 +192,10 @@ func Parse() (Config, error) {
}

if !slices.Contains([]string{ServiceDiscoveryProviderNomad, ServiceDiscoveryProviderKubernetes, ServiceDiscoveryProviderLocal}, config.ServiceDiscoveryProvider) {
return config, fmt.Errorf("invalid service discovery provider: %s", config.ServiceDiscoveryProvider)
return config, newFailureError(
FailureConditionInvalidServiceDiscoveryProvider,
fmt.Sprintf("invalid service discovery provider: %s", config.ServiceDiscoveryProvider),
)
}

return config, nil
Expand Down
11 changes: 11 additions & 0 deletions packages/api/internal/cfg/model_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,17 @@ func TestParse(t *testing.T) {
require.NoError(t, err)
assert.Equal(t, content, result.VolumesToken.SigningKey)
})

t.Run("invalid service discovery provider exposes failure condition", func(t *testing.T) {
t.Setenv("SERVICE_DISCOVERY_PROVIDER", "invalid")

_, err := Parse()
require.Error(t, err)

condition, ok := ParseFailureCondition(err)
require.True(t, ok)
assert.Equal(t, FailureConditionInvalidServiceDiscoveryProvider, condition)
})
}

// removeEnv was mostly copied from the implementation of t.Setenv
Expand Down
50 changes: 50 additions & 0 deletions packages/api/internal/handlers/store.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ import (
clickhouse "github.com/e2b-dev/infra/packages/clickhouse/pkg"
sqlcdb "github.com/e2b-dev/infra/packages/db/client"
authdb "github.com/e2b-dev/infra/packages/db/pkg/auth"
"github.com/e2b-dev/infra/packages/db/pkg/dberrors"
"github.com/e2b-dev/infra/packages/db/pkg/pool"
"github.com/e2b-dev/infra/packages/shared/pkg/apierrors"
"github.com/e2b-dev/infra/packages/shared/pkg/consts"
Expand Down Expand Up @@ -396,3 +397,52 @@ func (a *APIStore) GetTeamFromAuthProviderToken(ctx context.Context, ginCtx *gin

return a.authService.ValidateSupabaseTeam(ctx, ginCtx, teamID)
}

func (a *APIStore) GetTeamFromAdminToken(ctx context.Context, _ *gin.Context, teamID string) (*types.Team, *api.APIError) {
ctx, span := tracer.Start(ctx, "get team from admin token")
defer span.End()

teamUUID, err := uuid.Parse(teamID)
if err != nil {
return nil, &api.APIError{
Code: http.StatusBadRequest,
ClientMsg: "Invalid team ID",
Err: fmt.Errorf("failed to parse team ID: %w", err),
}
}

team, err := a.authService.GetTeamByID(ctx, teamUUID)
if err != nil {
var forbiddenErr *sharedauth.TeamForbiddenError
if errors.As(err, &forbiddenErr) {
return nil, &api.APIError{
Code: http.StatusForbidden,
ClientMsg: err.Error(),
Err: fmt.Errorf("failed getting team: %w", err),
}
}

if dberrors.IsNotFoundError(err) {
return nil, &api.APIError{
Code: http.StatusNotFound,
ClientMsg: "Team not found",
Err: fmt.Errorf("failed getting team: %w", err),
}
Comment thread
ben-fornefeld marked this conversation as resolved.
}

return nil, &api.APIError{
Code: http.StatusInternalServerError,
ClientMsg: "Backend authentication failed",
Err: fmt.Errorf("failed getting team: %w", err),
}
Comment thread
ben-fornefeld marked this conversation as resolved.
}
Comment thread
ben-fornefeld marked this conversation as resolved.
if team == nil {
return nil, &api.APIError{
Code: http.StatusNotFound,
ClientMsg: "Team not found",
Err: errors.New("team not found"),
}
}

return team, nil
}
10 changes: 8 additions & 2 deletions packages/api/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -197,7 +197,8 @@ func NewGinServer(ctx context.Context, config cfg.Config, tel *telemetry.Client,
auth.NewSupabaseTokenAuthenticator(apiStore.GetUserIDFromAuthProviderToken),
auth.NewSupabaseTeamAuthenticator(apiStore.GetTeamFromSupabaseToken),
auth.NewAuthProviderTeamAuthenticator(apiStore.GetTeamFromAuthProviderToken),
auth.NewAdminTokenAuthenticator(config.AdminToken),
auth.NewAdminApiKeyAuthenticator(config.AdminToken),
auth.NewAdminTeamAuthenticator(apiStore.GetTeamFromAdminToken),
},
metricsMiddleware.SetProcessingStartTime,
)
Expand Down Expand Up @@ -348,7 +349,12 @@ func run() int {

config, err := cfg.Parse()
if err != nil {
logger.L().Fatal(ctx, "Error parsing config", zap.Error(err))
fields := []zap.Field{zap.Error(err)}
if condition, ok := cfg.ParseFailureCondition(err); ok {
fields = append(fields, zap.String("config_failure_condition", string(condition)))
}

logger.L().Fatal(ctx, "Error parsing config", fields...)
}

err = sqlcdb.CheckMigrationVersion(ctx, config.PostgresConnectionString, expectedMigration)
Expand Down
30 changes: 27 additions & 3 deletions packages/auth/pkg/auth/middleware.go
Original file line number Diff line number Diff line change
Expand Up @@ -211,10 +211,14 @@ func NewAuthProviderTeamAuthenticator(validationFunc func(ctx context.Context, g
}
}

// NewAdminTokenAuthenticator creates an authenticator for the AdminTokenAuth security scheme (X-Admin-Token header).
func NewAdminTokenAuthenticator(adminToken string) Authenticator {
// NewAdminApiKeyAuthenticator creates an authenticator for the AdminApiKeyAuth security scheme (X-Admin-Token header).
func NewAdminApiKeyAuthenticator(adminToken string) Authenticator {
return newAdminApiKeyAuthenticator("AdminApiKeyAuth", adminToken)
}

func newAdminApiKeyAuthenticator(schemeName, adminToken string) Authenticator {
return &commonAuthenticator[struct{}]{
schemeName: "AdminTokenAuth",
schemeName: schemeName,
header: headerKey{
name: HeaderAdminToken,
},
Expand All @@ -223,6 +227,26 @@ func NewAdminTokenAuthenticator(adminToken string) Authenticator {
}
}

// NewAdminTeamAuthenticator creates an authenticator for AdminTeamAuth (X-Team-ID header).
func NewAdminTeamAuthenticator(validationFunc func(ctx context.Context, ginCtx *gin.Context, teamID string) (*types.Team, *APIError)) Authenticator {
return newAdminTeamAuthenticator("AdminTeamAuth", validationFunc)
}

func newAdminTeamAuthenticator(
schemeName string,
validationFunc func(ctx context.Context, ginCtx *gin.Context, teamID string) (*types.Team, *APIError),
) Authenticator {
return &commonAuthenticator[*types.Team]{
schemeName: schemeName,
header: headerKey{
name: HeaderTeamID,
},
validationFunc: validationFunc,
setContextFunc: setTeamInfo,
errorMessage: "Invalid admin token teamID.",
}
}

// CreateAuthenticationFunc creates an OpenAPI authentication function from a list of authenticators.
func CreateAuthenticationFunc(
authenticators []Authenticator,
Expand Down
52 changes: 52 additions & 0 deletions packages/auth/pkg/auth/middleware_test.go
Original file line number Diff line number Diff line change
@@ -1,9 +1,18 @@
package auth

import (
"context"
"net/http"
"net/http/httptest"
"testing"

"github.com/getkin/kin-openapi/openapi3filter"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"github.com/stretchr/testify/require"

"github.com/e2b-dev/infra/packages/auth/pkg/types"
authqueries "github.com/e2b-dev/infra/packages/db/pkg/auth/queries"
)

func TestAdminValidationFunction(t *testing.T) {
Expand All @@ -26,3 +35,46 @@ func TestAdminValidationFunction(t *testing.T) {
require.Equal(t, 401, err.Code)
})
}

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

ctx := context.Background()
teamID := uuid.New()
team := types.NewTeam(&authqueries.Team{ID: teamID}, &authqueries.TeamLimit{})

req := httptest.NewRequestWithContext(ctx, http.MethodGet, "/", nil)
req.Header.Set(HeaderTeamID, teamID.String())

ginCtx, _ := gin.CreateTestContext(httptest.NewRecorder())
authenticator := NewAdminTeamAuthenticator(func(_ context.Context, _ *gin.Context, gotTeamID string) (*types.Team, *APIError) {
if gotTeamID != teamID.String() {
return nil, &APIError{
Err: ErrInvalidAuthHeader,
ClientMsg: "Invalid team ID",
Code: http.StatusBadRequest,
}
}

return team, nil
})
if got, want := authenticator.SecuritySchemeName(), "AdminTeamAuth"; got != want {
t.Fatalf("NewAdminTeamAuthenticator().SecuritySchemeName() = %q, want %q", got, want)
}

err := authenticator.Authenticate(ctx, ginCtx, &openapi3filter.AuthenticationInput{
RequestValidationInput: &openapi3filter.RequestValidationInput{Request: req},
})
if err != nil {
t.Fatalf("AdminTeamAuth.Authenticate(valid team ID) error: %v", err)
}

got, ok := GetTeamInfo(ginCtx)
if !ok {
t.Fatalf("GetTeamInfo(ginCtx) ok = false, want true")
}

if got.Team.ID != teamID {
t.Errorf("GetTeamInfo(ginCtx).Team.ID = %s, want %s", got.Team.ID, teamID)
}
}
Loading
Loading