diff --git a/extension/contract/contract.go b/extension/contract/contract.go index 590c42fc..b1004625 100644 --- a/extension/contract/contract.go +++ b/extension/contract/contract.go @@ -333,6 +333,9 @@ func Register( if err := dispatcher.RegisterQuery(d, c, "auth.featureToggles", 1, featureTogglesHandler(deps)); err != nil { return fmt.Errorf("authsome/contract: register auth.featureToggles: %w", err) } + if err := dispatcher.RegisterQuery(d, c, "plugins.list", 1, pluginsListHandler(deps)); err != nil { + return fmt.Errorf("authsome/contract: register plugins.list: %w", err) + } if err := dispatcher.RegisterCommand(d, c, "auth.toggleFeature", 1, toggleFeatureHandler(deps)); err != nil { return fmt.Errorf("authsome/contract: register auth.toggleFeature: %w", err) } diff --git a/extension/contract/handlers_auth_pages.go b/extension/contract/handlers_auth_pages.go index a6035bee..15012b65 100644 --- a/extension/contract/handlers_auth_pages.go +++ b/extension/contract/handlers_auth_pages.go @@ -21,11 +21,17 @@ package contract import ( "context" "errors" + "net/url" + "regexp" "strings" + "sync" authsome "github.com/xraph/authsome" "github.com/xraph/authsome/account" + "github.com/xraph/authsome/app" + "github.com/xraph/authsome/environment" "github.com/xraph/authsome/formconfig" + "github.com/xraph/authsome/id" "github.com/xraph/authsome/user" dashauth "github.com/xraph/forge/extensions/dashboard/auth" @@ -211,16 +217,56 @@ func mapResetError(err error) error { // when no users exist yet — the dashboard's AuthGate uses this to // redirect to /setup before /login. type SetupStatusResponse struct { - Pending bool `json:"pending"` + Pending bool `json:"pending"` + Platform *SetupPlatformDefaults `json:"platform,omitempty"` + Environment *SetupEnvironmentDefaults `json:"environment,omitempty"` +} + +// SetupPlatformDefaults contains the public platform fields the anonymous +// setup form may edit. Keep this deliberately smaller than AppDetail: stored +// metadata and the publishable key must never cross this public boundary. +type SetupPlatformDefaults struct { + Name string `json:"name"` + Slug string `json:"slug"` + Logo string `json:"logo,omitempty"` +} + +// SetupEnvironmentDefaults contains only the public default-environment +// fields needed to initialize first-run setup. +type SetupEnvironmentDefaults struct { + Name string `json:"name"` + Slug string `json:"slug"` + Type string `json:"type"` + IsDefault bool `json:"isDefault"` + Color string `json:"color,omitempty"` + Description string `json:"description,omitempty"` } // SetupInput is the wire shape for auth.setup. Mirrors the // auth.setup-form renderer's fields. type SetupInput struct { - Email string `json:"email"` - Password string `json:"password"` - Name string `json:"name,omitempty"` - OrganizationName string `json:"organizationName,omitempty"` + Email string `json:"email"` + Password string `json:"password"` + Name string `json:"name,omitempty"` + OrganizationName string `json:"organizationName,omitempty"` + Platform *SetupPlatformInput `json:"platform,omitempty"` + Environment *SetupEnvironmentInput `json:"environment,omitempty"` +} + +type SetupPlatformInput struct { + Name string `json:"name"` + Slug string `json:"slug"` + Logo string `json:"logo,omitempty"` + Metadata map[string]string `json:"metadata,omitempty"` +} + +type SetupEnvironmentInput struct { + Name string `json:"name"` + Slug string `json:"slug"` + Type string `json:"type"` + Color string `json:"color,omitempty"` + Description string `json:"description,omitempty"` + Metadata map[string]string `json:"metadata,omitempty"` } // SetupResponse is the auth.setup reply. @@ -247,7 +293,36 @@ func setupStatusHandler(deps Deps) func(ctx context.Context, _ struct{}, _ contr // already-bootstrapped deployment whose count query just hiccupped. return SetupStatusResponse{Pending: false}, nil } - return SetupStatusResponse{Pending: list.Total == 0}, nil + if list.Total > 0 { + return SetupStatusResponse{Pending: false}, nil + } + + appID := defaultAppID(eng) + platform, err := eng.GetApp(ctx, appID) + if err != nil { + return SetupStatusResponse{}, mapEngineError(err) + } + defaultEnv, err := eng.GetDefaultEnvironment(ctx, appID) + if err != nil { + return SetupStatusResponse{}, mapEngineError(err) + } + + return SetupStatusResponse{ + Pending: true, + Platform: &SetupPlatformDefaults{ + Name: platform.Name, + Slug: platform.Slug, + Logo: platform.Logo, + }, + Environment: &SetupEnvironmentDefaults{ + Name: defaultEnv.Name, + Slug: defaultEnv.Slug, + Type: string(defaultEnv.Type), + IsDefault: defaultEnv.IsDefault, + Color: defaultEnv.Color, + Description: defaultEnv.Description, + }, + }, nil } } @@ -260,19 +335,199 @@ func setupStatusHandler(deps Deps) func(ctx context.Context, _ struct{}, _ contr // to create the org alongside the user. Failure to create the org // doesn't roll back the user — admins can create their org manually // post-setup. +const ( + setupNameMax = 80 + setupSlugMax = 63 + setupDescriptionMax = 240 + setupMetadataMax = 20 + setupMetadataKeyMax = 64 + setupMetadataValMax = 512 +) + +var setupSlugPattern = regexp.MustCompile(`^[a-z0-9]+(?:-[a-z0-9]+)*$`) + +func setupBadRequest(field, message string) error { + return &contract.Error{ + Code: contract.CodeBadRequest, + Message: message, + Details: map[string]any{"field": field}, + } +} + +func normalizeSetupMetadata(field string, metadata map[string]string) (map[string]string, error) { + if metadata == nil { + return nil, nil + } + if len(metadata) > setupMetadataMax { + return nil, setupBadRequest(field, "metadata may contain at most 20 entries") + } + normalized := make(map[string]string, len(metadata)) + for rawKey, rawValue := range metadata { + key := strings.TrimSpace(rawKey) + value := strings.TrimSpace(rawValue) + if key == "" || value == "" { + return nil, setupBadRequest(field, "metadata keys and values are required") + } + if len(key) > setupMetadataKeyMax { + return nil, setupBadRequest(field, "metadata keys may contain at most 64 characters") + } + if len(value) > setupMetadataValMax { + return nil, setupBadRequest(field, "metadata values may contain at most 512 characters") + } + if _, exists := normalized[key]; exists { + return nil, setupBadRequest(field, "metadata keys must be unique") + } + normalized[key] = value + } + return normalized, nil +} + +func validateSetupNameAndSlug(prefix, name, slug string) (trimmedName, trimmedSlug string, err error) { + name = strings.TrimSpace(name) + slug = strings.TrimSpace(slug) + if name == "" { + return "", "", setupBadRequest(prefix+".name", "name is required") + } + if len(name) > setupNameMax { + return "", "", setupBadRequest(prefix+".name", "name may contain at most 80 characters") + } + if slug == "" { + return "", "", setupBadRequest(prefix+".slug", "slug is required") + } + if len(slug) > setupSlugMax || !setupSlugPattern.MatchString(slug) { + return "", "", setupBadRequest(prefix+".slug", "slug must use lowercase letters, numbers, and single hyphens") + } + return name, slug, nil +} + +func validSetupLogo(value string) bool { + if value == "" || (strings.HasPrefix(value, "/") && !strings.HasPrefix(value, "//")) { + return true + } + parsed, err := url.Parse(value) + return err == nil && parsed.Host != "" && (parsed.Scheme == "http" || parsed.Scheme == "https") +} + +func validateSetupInput(in SetupInput) (SetupInput, error) { + in.Email = strings.ToLower(strings.TrimSpace(in.Email)) + in.Name = strings.TrimSpace(in.Name) + in.OrganizationName = strings.TrimSpace(in.OrganizationName) + if in.Email == "" || in.Password == "" { + return SetupInput{}, setupBadRequest("administrator.email", "email and password are required") + } + + if in.Platform != nil { + platform := *in.Platform + var err error + platform.Name, platform.Slug, err = validateSetupNameAndSlug("platform", platform.Name, platform.Slug) + if err != nil { + return SetupInput{}, err + } + platform.Logo = strings.TrimSpace(platform.Logo) + if !validSetupLogo(platform.Logo) { + return SetupInput{}, setupBadRequest("platform.logo", "logo must be an HTTP URL or a root-relative path") + } + platform.Metadata, err = normalizeSetupMetadata("platform.metadata", platform.Metadata) + if err != nil { + return SetupInput{}, err + } + in.Platform = &platform + } + + if in.Environment != nil { + env := *in.Environment + var err error + env.Name, env.Slug, err = validateSetupNameAndSlug("environment", env.Name, env.Slug) + if err != nil { + return SetupInput{}, err + } + env.Type = strings.TrimSpace(env.Type) + if !environment.Type(env.Type).IsValid() { + return SetupInput{}, setupBadRequest("environment.type", "environment type is invalid") + } + env.Color = strings.TrimSpace(env.Color) + env.Description = strings.TrimSpace(env.Description) + if len(env.Description) > setupDescriptionMax { + return SetupInput{}, setupBadRequest("environment.description", "description may contain at most 240 characters") + } + env.Metadata, err = normalizeSetupMetadata("environment.metadata", env.Metadata) + if err != nil { + return SetupInput{}, err + } + in.Environment = &env + } + + return in, nil +} + +func mergeSetupMetadata[T ~map[string]string](existing T, incoming map[string]string) T { + if incoming == nil { + return existing + } + merged := make(T, len(existing)+len(incoming)) + for key, value := range existing { + merged[key] = value + } + for key, value := range incoming { + merged[key] = value + } + return merged +} + +// applySetupPlatform writes the validated platform block onto the +// bootstrapped platform app. A nil block leaves the app untouched. +func applySetupPlatform(ctx context.Context, eng *authsome.Engine, appID id.AppID, in *SetupPlatformInput) error { + if in == nil { + return nil + } + current, err := eng.GetApp(ctx, appID) + if err != nil { + return err + } + current.Name = in.Name + current.Slug = in.Slug + current.Logo = in.Logo + current.Metadata = mergeSetupMetadata[app.Metadata](current.Metadata, in.Metadata) + return eng.UpdateApp(ctx, current) +} + +// applySetupEnvironment writes the validated environment block onto the +// app's default environment. Setup never creates or removes environments. +func applySetupEnvironment(ctx context.Context, eng *authsome.Engine, appID id.AppID, in *SetupEnvironmentInput) error { + if in == nil { + return nil + } + current, err := eng.GetDefaultEnvironment(ctx, appID) + if err != nil { + return err + } + current.Name = in.Name + current.Slug = in.Slug + current.Type = environment.Type(in.Type) + current.Color = in.Color + current.Description = in.Description + current.Metadata = mergeSetupMetadata[environment.Metadata](current.Metadata, in.Metadata) + return eng.UpdateEnvironment(ctx, current) +} + func setupHandler(deps Deps) func(ctx context.Context, in SetupInput, _ contract.Principal) (SetupResponse, error) { + var setupMu sync.Mutex return func(ctx context.Context, in SetupInput, _ contract.Principal) (SetupResponse, error) { eng := deps.Engine if eng == nil { return SetupResponse{}, &contract.Error{Code: contract.CodeUnavailable, Message: "auth engine not configured"} } - email := strings.ToLower(strings.TrimSpace(in.Email)) - if email == "" || in.Password == "" { - return SetupResponse{}, &contract.Error{Code: contract.CodeBadRequest, Message: "email and password are required"} + normalized, err := validateSetupInput(in) + if err != nil { + return SetupResponse{}, err } + setupMu.Lock() + defer setupMu.Unlock() + // Refuse to bootstrap a populated deployment. - list, err := eng.AdminListUsers(ctx, &user.Query{AppID: defaultAppID(eng), Limit: 1}) + appID := defaultAppID(eng) + list, err := eng.AdminListUsers(ctx, &user.Query{AppID: appID, Limit: 1}) if err != nil { return SetupResponse{}, mapEngineError(err) } @@ -286,11 +541,18 @@ func setupHandler(deps Deps) func(ctx context.Context, in SetupInput, _ contract return SetupResponse{}, &contract.Error{Code: contract.CodeInternal, Message: "no http context (forge >= dashauth.WithHTTP required)"} } - first, last := splitName(in.Name) + if err = applySetupPlatform(ctx, eng, appID, normalized.Platform); err != nil { + return SetupResponse{}, mapEngineError(err) + } + if err = applySetupEnvironment(ctx, eng, appID, normalized.Environment); err != nil { + return SetupResponse{}, mapEngineError(err) + } + + first, last := splitName(normalized.Name) req := &account.SignUpRequest{ - AppID: defaultAppID(eng), - Email: email, - Password: in.Password, + AppID: appID, + Email: normalized.Email, + Password: normalized.Password, FirstName: first, LastName: last, IPAddress: clientIP(httpReq), @@ -308,7 +570,7 @@ func setupHandler(deps Deps) func(ctx context.Context, in SetupInput, _ contract // surface — we don't know at compile time whether the org // plugin is loaded. Skip it for now and let admins create their // org through the dashboard's /organizations page post-setup. - _ = in.OrganizationName + _ = normalized.OrganizationName subject := "" if u != nil { diff --git a/extension/contract/handlers_auth_pages_test.go b/extension/contract/handlers_auth_pages_test.go new file mode 100644 index 00000000..cfdc0b5a --- /dev/null +++ b/extension/contract/handlers_auth_pages_test.go @@ -0,0 +1,412 @@ +package contract + +import ( + "context" + "errors" + "fmt" + "net/http/httptest" + "reflect" + "strings" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/require" + + authsome "github.com/xraph/authsome" + "github.com/xraph/authsome/account" + "github.com/xraph/authsome/app" + "github.com/xraph/authsome/environment" + "github.com/xraph/authsome/id" + "github.com/xraph/authsome/internal/secutil" + "github.com/xraph/authsome/rbac" + "github.com/xraph/authsome/store" + "github.com/xraph/authsome/store/memory" + "github.com/xraph/authsome/user" + + dashcontract "github.com/xraph/forge/extensions/dashboard/contract" + "github.com/xraph/warden" + wardenmem "github.com/xraph/warden/store/memory" + "golang.org/x/crypto/bcrypt" +) + +type failingUserListStore struct { + store.Store + fail atomic.Bool +} + +func (s *failingUserListStore) ListUsers(ctx context.Context, q *user.Query) (*user.List, error) { + if s.fail.Load() { + return nil, fmt.Errorf("forced user-list failure") + } + return s.Store.ListUsers(ctx, q) +} + +type failingEnvironmentUpdateStore struct { + store.Store + fail atomic.Bool +} + +func (s *failingEnvironmentUpdateStore) UpdateEnvironment(ctx context.Context, env *environment.Environment) error { + if s.fail.Swap(false) { + return fmt.Errorf("forced environment-update failure") + } + return s.Store.UpdateEnvironment(ctx, env) +} + +func newSetupEngine(t *testing.T) *authsome.Engine { + t.Helper() + cfg := authsome.DefaultConfig() + cfg.Password.BcryptCost = bcrypt.MinCost + return secutil.NewTestEngine(t, + authsome.WithConfig(cfg), + authsome.WithBootstrap(), + ) +} + +func newSetupEngineWithFailingUserList(t *testing.T) (*authsome.Engine, *failingUserListStore) { + t.Helper() + wrapped := &failingUserListStore{Store: memory.New()} + eng := startSetupEngineWithStore(t, wrapped) + return eng, wrapped +} + +func startSetupEngineWithStore(t *testing.T, setupStore store.Store) *authsome.Engine { + t.Helper() + w, err := warden.NewEngine(warden.WithStore(wardenmem.New())) + require.NoError(t, err) + cfg := authsome.DefaultConfig() + cfg.Password.BcryptCost = bcrypt.MinCost + eng, err := authsome.NewEngine( + authsome.WithStore(setupStore), + authsome.WithWarden(w), + authsome.WithDisableMigrate(), + authsome.WithConfig(cfg), + authsome.WithBootstrap(), + ) + require.NoError(t, err) + require.NoError(t, eng.Start(context.Background())) + t.Cleanup(func() { _ = eng.Stop(context.Background()) }) + secutil.RelaxAuthDefaults(t, eng) + return eng +} + +func TestSetupStatusReturnsSafeBootstrapDefaults(t *testing.T) { + eng := newSetupEngine(t) + got, err := setupStatusHandler(Deps{Engine: eng})( + context.Background(), struct{}{}, dashcontract.Principal{}, + ) + require.NoError(t, err) + require.True(t, got.Pending) + require.NotNil(t, got.Platform) + require.Equal(t, "Platform", got.Platform.Name) + require.Equal(t, "platform", got.Platform.Slug) + require.NotNil(t, got.Environment) + require.True(t, got.Environment.IsDefault) + require.Equal(t, "development", got.Environment.Type) +} + +func TestSetupStatusDefaultsCannotExposePrivateFields(t *testing.T) { + for _, tc := range []struct { + name string + typ reflect.Type + }{ + {name: "platform", typ: reflect.TypeOf(SetupPlatformDefaults{})}, + {name: "environment", typ: reflect.TypeOf(SetupEnvironmentDefaults{})}, + } { + t.Run(tc.name, func(t *testing.T) { + for _, field := range []string{"Metadata", "PublishableKey", "Settings", "Credentials"} { + _, found := tc.typ.FieldByName(field) + require.False(t, found, "%s response must not expose %s", tc.name, field) + } + }) + } +} + +func TestSetupStatusUnavailableWithoutEngine(t *testing.T) { + _, err := setupStatusHandler(Deps{})( + context.Background(), struct{}{}, dashcontract.Principal{}, + ) + var contractErr *dashcontract.Error + require.True(t, errors.As(err, &contractErr)) + require.Equal(t, dashcontract.CodeUnavailable, contractErr.Code) +} + +func TestSetupStatusFailsClosedWhenUserCountFails(t *testing.T) { + eng, wrapped := newSetupEngineWithFailingUserList(t) + wrapped.fail.Store(true) + + got, err := setupStatusHandler(Deps{Engine: eng})( + context.Background(), struct{}{}, dashcontract.Principal{}, + ) + require.NoError(t, err) + require.False(t, got.Pending) + require.Nil(t, got.Platform) + require.Nil(t, got.Environment) +} + +func TestSetupStatusOmitsDefaultsAfterFirstUser(t *testing.T) { + eng := newSetupEngine(t) + _, _, err := eng.SignUp(context.Background(), &account.SignUpRequest{ + AppID: eng.PlatformAppID(), + Email: "owner@example.com", + Password: "SecureP@ss1", + }) + require.NoError(t, err) + + got, err := setupStatusHandler(Deps{Engine: eng})( + context.Background(), struct{}{}, dashcontract.Principal{}, + ) + require.NoError(t, err) + require.Equal(t, SetupStatusResponse{Pending: false}, got) +} + +func completeSetupInput() SetupInput { + return SetupInput{ + Email: "owner@example.com", + Password: "SecureP@ss1", + Name: "Ada Lovelace", + Platform: &SetupPlatformInput{ + Name: "TwinOS Office", + Slug: "twinos-office", + Logo: "https://example.test/forge.svg", + Metadata: map[string]string{"region": "us-central"}, + }, + Environment: &SetupEnvironmentInput{ + Name: "Local Development", + Slug: "local-development", + Type: "development", + Color: "#2563eb", + Description: "Local plugin development", + Metadata: map[string]string{"purpose": "plugins"}, + }, + } +} + +func TestSetupHandlerAppliesPlatformEnvironmentAndOwner(t *testing.T) { + eng := newSetupEngine(t) + ctx := context.Background() + + platform, err := eng.GetApp(ctx, eng.PlatformAppID()) + require.NoError(t, err) + platform.Metadata = app.Metadata{"existing": "kept"} + require.NoError(t, eng.UpdateApp(ctx, platform)) + defaultEnv, err := eng.GetDefaultEnvironment(ctx, eng.PlatformAppID()) + require.NoError(t, err) + defaultEnv.Metadata = environment.Metadata{"seeded": "kept"} + require.NoError(t, eng.UpdateEnvironment(ctx, defaultEnv)) + + h := setupHandler(Deps{Engine: eng}) + httpCtx, writer, _ := withHTTPCtx(t) + got, err := h(httpCtx, completeSetupInput(), dashcontract.Principal{}) + require.NoError(t, err) + require.True(t, got.OK) + require.NotEmpty(t, got.Subject) + + updatedApp, err := eng.GetApp(ctx, eng.PlatformAppID()) + require.NoError(t, err) + require.Equal(t, "TwinOS Office", updatedApp.Name) + require.Equal(t, "twinos-office", updatedApp.Slug) + require.Equal(t, "https://example.test/forge.svg", updatedApp.Logo) + require.Equal(t, "kept", updatedApp.Metadata["existing"]) + require.Equal(t, "us-central", updatedApp.Metadata["region"]) + + updatedEnv, err := eng.GetDefaultEnvironment(ctx, eng.PlatformAppID()) + require.NoError(t, err) + require.Equal(t, defaultEnv.ID, updatedEnv.ID) + require.Equal(t, "Local Development", updatedEnv.Name) + require.Equal(t, "local-development", updatedEnv.Slug) + require.Equal(t, environment.TypeDevelopment, updatedEnv.Type) + require.Equal(t, "#2563eb", updatedEnv.Color) + require.Equal(t, "Local plugin development", updatedEnv.Description) + require.Equal(t, "kept", updatedEnv.Metadata["seeded"]) + require.Equal(t, "plugins", updatedEnv.Metadata["purpose"]) + + cookies := writer.(*httptest.ResponseRecorder).Result().Cookies() + require.Condition(t, func() bool { + for _, cookie := range cookies { + if cookie.Name == dashboardCookieName && cookie.Value != "" { + return true + } + } + return false + }, "setup must write the dashboard session cookie") + + userID, err := id.ParseUserID(got.Subject) + require.NoError(t, err) + roles, err := eng.ListUserRoles(ctx, userID) + require.NoError(t, err) + require.Condition(t, func() bool { + for _, role := range roles { + if role.Slug == rbac.PlatformOwnerSlug { + return true + } + } + return false + }, "first setup user must receive platform-owner") + + httpCtx, _, _ = withHTTPCtx(t) + _, err = h(httpCtx, completeSetupInput(), dashcontract.Principal{}) + var contractErr *dashcontract.Error + require.True(t, errors.As(err, &contractErr)) + require.Equal(t, dashcontract.CodePermissionDenied, contractErr.Code) +} + +func TestSetupHandlerLegacyPayloadPreservesBootstrapConfiguration(t *testing.T) { + eng := newSetupEngine(t) + ctx := context.Background() + beforeApp, err := eng.GetApp(ctx, eng.PlatformAppID()) + require.NoError(t, err) + beforeEnv, err := eng.GetDefaultEnvironment(ctx, eng.PlatformAppID()) + require.NoError(t, err) + + httpCtx, _, _ := withHTTPCtx(t) + got, err := setupHandler(Deps{Engine: eng})(httpCtx, SetupInput{ + Email: "legacy@example.com", Password: "SecureP@ss1", + }, dashcontract.Principal{}) + require.NoError(t, err) + require.True(t, got.OK) + + afterApp, err := eng.GetApp(ctx, eng.PlatformAppID()) + require.NoError(t, err) + afterEnv, err := eng.GetDefaultEnvironment(ctx, eng.PlatformAppID()) + require.NoError(t, err) + require.Equal(t, beforeApp.Name, afterApp.Name) + require.Equal(t, beforeApp.Slug, afterApp.Slug) + require.Equal(t, beforeEnv.ID, afterEnv.ID) + require.Equal(t, beforeEnv.Name, afterEnv.Name) + require.Equal(t, beforeEnv.Slug, afterEnv.Slug) +} + +func TestSetupHandlerRejectsInvalidConfigurationBeforeWriting(t *testing.T) { + metadataOverLimit := make(map[string]string, 21) + for i := 0; i < 21; i++ { + metadataOverLimit[fmt.Sprintf("key-%d", i)] = "value" + } + + tests := []struct { + name string + mutate func(*SetupInput) + field string + }{ + {name: "blank platform name", field: "platform.name", mutate: func(in *SetupInput) { in.Platform.Name = " " }}, + {name: "malformed platform slug", field: "platform.slug", mutate: func(in *SetupInput) { in.Platform.Slug = "TwinOS Office" }}, + {name: "unsupported environment type", field: "environment.type", mutate: func(in *SetupInput) { in.Environment.Type = "preview" }}, + {name: "too many metadata entries", field: "platform.metadata", mutate: func(in *SetupInput) { in.Platform.Metadata = metadataOverLimit }}, + {name: "blank metadata key", field: "platform.metadata", mutate: func(in *SetupInput) { in.Platform.Metadata = map[string]string{" ": "value"} }}, + {name: "blank metadata value", field: "environment.metadata", mutate: func(in *SetupInput) { in.Environment.Metadata = map[string]string{"key": " "} }}, + {name: "duplicate trimmed metadata key", field: "platform.metadata", mutate: func(in *SetupInput) { in.Platform.Metadata = map[string]string{" key": "one", "key ": "two"} }}, + {name: "oversized metadata key", field: "platform.metadata", mutate: func(in *SetupInput) { in.Platform.Metadata = map[string]string{strings.Repeat("k", 65): "value"} }}, + {name: "oversized metadata value", field: "environment.metadata", mutate: func(in *SetupInput) { in.Environment.Metadata = map[string]string{"key": strings.Repeat("v", 513)} }}, + {name: "invalid logo scheme", field: "platform.logo", mutate: func(in *SetupInput) { in.Platform.Logo = "javascript:alert(1)" }}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + eng := newSetupEngine(t) + ctx := context.Background() + beforeApp, err := eng.GetApp(ctx, eng.PlatformAppID()) + require.NoError(t, err) + beforeEnv, err := eng.GetDefaultEnvironment(ctx, eng.PlatformAppID()) + require.NoError(t, err) + + in := completeSetupInput() + tc.mutate(&in) + httpCtx, _, _ := withHTTPCtx(t) + _, err = setupHandler(Deps{Engine: eng})(httpCtx, in, dashcontract.Principal{}) + var contractErr *dashcontract.Error + require.True(t, errors.As(err, &contractErr)) + require.Equal(t, dashcontract.CodeBadRequest, contractErr.Code) + require.Equal(t, tc.field, contractErr.Details["field"]) + + afterApp, appErr := eng.GetApp(ctx, eng.PlatformAppID()) + require.NoError(t, appErr) + afterEnv, envErr := eng.GetDefaultEnvironment(ctx, eng.PlatformAppID()) + require.NoError(t, envErr) + require.Equal(t, beforeApp.Name, afterApp.Name) + require.Equal(t, beforeApp.Slug, afterApp.Slug) + require.Equal(t, beforeEnv.ID, afterEnv.ID) + require.Equal(t, beforeEnv.Name, afterEnv.Name) + users, listErr := eng.AdminListUsers(ctx, &user.Query{AppID: eng.PlatformAppID(), Limit: 1}) + require.NoError(t, listErr) + require.Zero(t, users.Total) + }) + } +} + +func TestSetupHandlerRetriesAfterEnvironmentUpdateFailure(t *testing.T) { + baseStore := memory.New() + wrapped := &failingEnvironmentUpdateStore{Store: baseStore} + eng := startSetupEngineWithStore(t, wrapped) + ctx := context.Background() + beforeEnv, err := eng.GetDefaultEnvironment(ctx, eng.PlatformAppID()) + require.NoError(t, err) + h := setupHandler(Deps{Engine: eng}) + + wrapped.fail.Store(true) + httpCtx, _, _ := withHTTPCtx(t) + _, err = h(httpCtx, completeSetupInput(), dashcontract.Principal{}) + require.Error(t, err) + + updatedApp, err := eng.GetApp(ctx, eng.PlatformAppID()) + require.NoError(t, err) + require.Equal(t, "TwinOS Office", updatedApp.Name) + users, err := eng.AdminListUsers(ctx, &user.Query{AppID: eng.PlatformAppID(), Limit: 1}) + require.NoError(t, err) + require.Zero(t, users.Total) + + httpCtx, _, _ = withHTTPCtx(t) + got, err := h(httpCtx, completeSetupInput(), dashcontract.Principal{}) + require.NoError(t, err) + require.True(t, got.OK) + afterEnv, err := eng.GetDefaultEnvironment(ctx, eng.PlatformAppID()) + require.NoError(t, err) + require.Equal(t, beforeEnv.ID, afterEnv.ID) + apps, err := eng.ListApps(ctx) + require.NoError(t, err) + require.Len(t, apps, 1) + users, err = eng.AdminListUsers(ctx, &user.Query{AppID: eng.PlatformAppID(), Limit: 2}) + require.NoError(t, err) + require.Equal(t, 1, users.Total) +} + +func TestSetupHandlerConcurrentCreatesOneOwner(t *testing.T) { + eng := newSetupEngine(t) + h := setupHandler(Deps{Engine: eng}) + start := make(chan struct{}) + type result struct { + response SetupResponse + err error + } + results := make(chan result, 2) + + for i := 0; i < 2; i++ { + in := completeSetupInput() + in.Email = fmt.Sprintf("owner-%d@example.com", i) + go func() { + <-start + httpCtx, _, _ := withHTTPCtx(t) + response, err := h(httpCtx, in, dashcontract.Principal{}) + results <- result{response: response, err: err} + }() + } + close(start) + + successes := 0 + permissionDenied := 0 + for i := 0; i < 2; i++ { + got := <-results + if got.err == nil && got.response.OK { + successes++ + continue + } + var contractErr *dashcontract.Error + if errors.As(got.err, &contractErr) && contractErr.Code == dashcontract.CodePermissionDenied { + permissionDenied++ + } + } + require.Equal(t, 1, successes) + require.Equal(t, 1, permissionDenied) + users, err := eng.AdminListUsers(context.Background(), &user.Query{AppID: eng.PlatformAppID(), Limit: 2}) + require.NoError(t, err) + require.Equal(t, 1, users.Total) +} diff --git a/extension/contract/handlers_plugins.go b/extension/contract/handlers_plugins.go new file mode 100644 index 00000000..ddc9452a --- /dev/null +++ b/extension/contract/handlers_plugins.go @@ -0,0 +1,38 @@ +package contract + +import ( + "context" + "sort" + + "github.com/xraph/forge/extensions/dashboard/contract" +) + +// PluginSummary describes an installed Authsome plugin. Settings counts are +// taken from the live settings registry, not the dashboard contributors. +type PluginSummary struct { + Name string `json:"name"` + SettingCount int `json:"settingCount"` +} + +type PluginsListResponse struct { + Plugins []PluginSummary `json:"plugins"` +} + +func pluginsListHandler(deps Deps) func(context.Context, struct{}, contract.Principal) (PluginsListResponse, error) { + return func(_ context.Context, _ struct{}, _ contract.Principal) (PluginsListResponse, error) { + if deps.Engine == nil || deps.Engine.Plugins() == nil { + return PluginsListResponse{}, &contract.Error{Code: contract.CodeUnavailable, Message: "auth plugin registry not configured"} + } + out := PluginsListResponse{Plugins: make([]PluginSummary, 0, len(deps.Engine.Plugins().Plugins()))} + for _, installed := range deps.Engine.Plugins().Plugins() { + name := installed.Name() + summary := PluginSummary{Name: name} + if mgr := deps.Engine.Settings(); mgr != nil { + summary.SettingCount = len(mgr.DefinitionsForNamespace(name)) + } + out.Plugins = append(out.Plugins, summary) + } + sort.Slice(out.Plugins, func(i, j int) bool { return out.Plugins[i].Name < out.Plugins[j].Name }) + return out, nil + } +} diff --git a/extension/contract/handlers_plugins_test.go b/extension/contract/handlers_plugins_test.go new file mode 100644 index 00000000..4ce8c92c --- /dev/null +++ b/extension/contract/handlers_plugins_test.go @@ -0,0 +1,40 @@ +package contract + +import ( + "context" + "testing" + + "github.com/xraph/forge/extensions/dashboard/contract" + "github.com/xraph/warden" + wardenmem "github.com/xraph/warden/store/memory" + + authsome "github.com/xraph/authsome" + "github.com/xraph/authsome/store/memory" +) + +type inventoryPlugin string + +func (p inventoryPlugin) Name() string { return string(p) } + +func TestPluginsListHandler(t *testing.T) { + w, err := warden.NewEngine(warden.WithStore(wardenmem.New())) + if err != nil { + t.Fatal(err) + } + engine, err := authsome.NewEngine( + authsome.WithStore(memory.New()), + authsome.WithWarden(w), + authsome.WithPlugin(inventoryPlugin("zeta")), + authsome.WithPlugin(inventoryPlugin("alpha")), + ) + if err != nil { + t.Fatal(err) + } + result, err := pluginsListHandler(Deps{Engine: engine})(context.Background(), struct{}{}, contract.Principal{}) + if err != nil { + t.Fatal(err) + } + if len(result.Plugins) != 2 || result.Plugins[0].Name != "alpha" || result.Plugins[1].Name != "zeta" { + t.Fatalf("unexpected installed plugins: %+v", result.Plugins) + } +} diff --git a/extension/contract/manifest.yaml b/extension/contract/manifest.yaml index eee7c1e6..41b22d61 100644 --- a/extension/contract/manifest.yaml +++ b/extension/contract/manifest.yaml @@ -129,10 +129,11 @@ intents: - { name: auth.forgotPassword, kind: command, version: 1, capability: write } - { name: auth.resetPassword, kind: command, version: 1, capability: write } - { name: auth.setupStatus, kind: query, version: 1, capability: read } - - { name: auth.setup, kind: command, version: 1, capability: write } + - { name: auth.setup, kind: command, version: 1, capability: write, invalidates: [auth.config, auth.setupStatus, apps.context, apps.list, environments.list] } - { name: auth.dynamicConfig, kind: query, version: 1, capability: read } - { name: auth.dynamicRegister, kind: command, version: 1, capability: write } - { name: auth.featureToggles, kind: query, version: 1, capability: read } + - { name: plugins.list, kind: query, version: 1, capability: read } - { name: auth.toggleFeature, kind: command, version: 1, capability: write, invalidates: [auth.featureToggles] } # Phase C.17 — App + environment switcher. apps.context drives the diff --git a/extension/contract/manifest_test.go b/extension/contract/manifest_test.go index 984ddc5b..0d245755 100644 --- a/extension/contract/manifest_test.go +++ b/extension/contract/manifest_test.go @@ -16,17 +16,10 @@ func TestManifest_Loads(t *testing.T) { if m.Contributor.Name != "auth" { t.Errorf("contributor name = %q, want auth", m.Contributor.Name) } - // 68 intents: 66 prior + 2 new feature-toggle intents - // (auth.featureToggles, auth.toggleFeature). apikeys.* are owned + // Includes the installed-plugin inventory query. apikeys.* are owned // by the apikey plugin manifest, not declared here. - if got := len(m.Intents); got != 68 { - t.Errorf("intents = %d, want 68 (with feature toggles)", got) - } - // 28 top-level graph routes: 32 prior - 4 routes that moved to - // their owning plugins (/organizations, /organizations/:id, /plans, - // /plans/:id). - if got := len(m.Graph); got != 28 { - t.Errorf("graph routes = %d, want 28 (org + plan pages moved to plugins)", got) + if got := len(m.Intents); got != 69 { + t.Errorf("intents = %d, want 69", got) } } @@ -49,18 +42,13 @@ func TestManifest_RegistersWithRegistry(t *testing.T) { if err := reg.Register(m); err != nil { t.Fatalf("register: %v", err) } - // Sanity-check the /login graph route survived registration. Slice - // (l.5) shifted the route from a hardcoded form.edit to the dynamic - // auth.login.form intent backed by the auth.config query, so the - // expectation flips to verifying the data binding. - root, ok := reg.MergedGraph("auth", "/login") - if !ok { - t.Fatal("expected /login route to be registered") - } - if root.Intent != "auth.login.form" { - t.Errorf("unexpected /login root: intent=%s", root.Intent) - } - if root.Data == nil || root.Data.QueryRef != "queries.config" { - t.Errorf("expected data: queries.config, got %+v", root.Data) + for _, name := range []string{"auth.config", "plugins.list"} { + intent, ok := reg.Intent("auth", name, 1) + if !ok { + t.Fatalf("expected %s to be registered", name) + } + if intent.Kind != "query" { + t.Errorf("%s kind = %q, want query", name, intent.Kind) + } } } diff --git a/extension/contract_catalog.go b/extension/contract_catalog.go new file mode 100644 index 00000000..0117829d --- /dev/null +++ b/extension/contract_catalog.go @@ -0,0 +1,71 @@ +package extension + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + + contract "github.com/xraph/forge/extensions/dashboard/contract" + "github.com/xraph/forge/extensions/dashboard/contract/remote" +) + +func fetchContractCatalog(ctx context.Context, baseURL, apiKey string) ([]*contract.ContractManifest, error) { + request, err := http.NewRequestWithContext(ctx, http.MethodGet, strings.TrimRight(baseURL, "/")+remote.DefaultManifestPath, http.NoBody) + if err != nil { + return nil, fmt.Errorf("build manifest request: %w", err) + } + request.Header.Set("Accept", "application/json") + if apiKey != "" { + request.Header.Set("Authorization", "Bearer "+apiKey) + } + client := &http.Client{Timeout: remote.DefaultTimeout} + response, err := client.Do(request) + if err != nil { + return nil, fmt.Errorf("fetch manifest catalog: %w", err) + } + defer response.Body.Close() + if response.StatusCode < 200 || response.StatusCode >= 300 { + return nil, fmt.Errorf("manifest endpoint returned HTTP %d", response.StatusCode) + } + const maxCatalogBytes = 4 << 20 + body, err := io.ReadAll(io.LimitReader(response.Body, maxCatalogBytes+1)) + if err != nil { + return nil, fmt.Errorf("read manifest catalog: %w", err) + } + if len(body) > maxCatalogBytes { + return nil, fmt.Errorf("manifest catalog exceeds %d bytes", maxCatalogBytes) + } + var catalog struct { + Manifests []*contract.ContractManifest `json:"manifests"` + } + if err := json.Unmarshal(body, &catalog); err != nil { + return nil, fmt.Errorf("decode manifest catalog: %w", err) + } + if catalog.Manifests == nil { + var manifest contract.ContractManifest + if err := json.Unmarshal(body, &manifest); err != nil { + return nil, fmt.Errorf("decode manifest: %w", err) + } + catalog.Manifests = []*contract.ContractManifest{&manifest} + } + if len(catalog.Manifests) == 0 { + return nil, fmt.Errorf("manifest catalog is empty") + } + seen := make(map[string]bool, len(catalog.Manifests)) + for _, manifest := range catalog.Manifests { + if manifest == nil || manifest.Contributor.Name == "" { + return nil, fmt.Errorf("manifest is missing contributor.name") + } + if seen[manifest.Contributor.Name] { + return nil, fmt.Errorf("duplicate contributor %q", manifest.Contributor.Name) + } + seen[manifest.Contributor.Name] = true + } + if !seen["auth"] { + return nil, fmt.Errorf("manifest catalog is missing auth") + } + return catalog.Manifests, nil +} diff --git a/extension/contract_catalog_test.go b/extension/contract_catalog_test.go new file mode 100644 index 00000000..8ac97347 --- /dev/null +++ b/extension/contract_catalog_test.go @@ -0,0 +1,89 @@ +package extension + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/xraph/forge" + contract "github.com/xraph/forge/extensions/dashboard/contract" + "github.com/xraph/forge/extensions/dashboard/contract/dispatcher" + "github.com/xraph/warden" + wardenmem "github.com/xraph/warden/store/memory" + + authsome "github.com/xraph/authsome" + "github.com/xraph/authsome/plugins/apikey" + "github.com/xraph/authsome/store/memory" +) + +func TestContractServerExportsInstalledPlugins(t *testing.T) { + wardenEngine, err := warden.NewEngine(warden.WithStore(wardenmem.New())) + if err != nil { + t.Fatal(err) + } + engine, err := authsome.NewEngine(authsome.WithStore(memory.New()), authsome.WithWarden(wardenEngine), authsome.WithPlugin(apikey.New())) + if err != nil { + t.Fatal(err) + } + extension := New() + extension.engine = engine + extension.SetLogger(forge.NewNoopLogger()) + router := forge.NewRouter() + if err := extension.registerContractServer(router); err != nil { + t.Fatal(err) + } + if _, ok := extension.contractReg.Intent("apikey", "apikeys.list", 1); !ok { + t.Fatal("installed API key contributor missing from contract server") + } +} + +func TestRemoteContractCatalogRegistersAndDispatchesPlugin(t *testing.T) { + manifest := func(name, intent string) *contract.ContractManifest { + return &contract.ContractManifest{SchemaVersion: 1, Contributor: contract.Contributor{Name: name, Envelope: contract.EnvelopeSupport{Supports: []string{"v1"}, Preferred: "v1"}}, Intents: []contract.Intent{{Name: intent, Kind: "query", Version: 1, Capability: "read"}}} + } + upstream := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + if request.URL.Path == "/authsome/_forge/contract/manifest" { + if request.Header.Get("Authorization") != "Bearer service-key" { + t.Error("missing service authorization") + } + _ = json.NewEncoder(writer).Encode(map[string]any{"manifests": []*contract.ContractManifest{manifest("auth", "auth.config"), manifest("apikey", "apikeys.list")}}) + return + } + if request.URL.Path != "/authsome/_forge/contract/dispatch" { + http.NotFound(writer, request) + return + } + var envelope contract.Request + if err := json.NewDecoder(request.Body).Decode(&envelope); err != nil { + t.Error(err) + } + if envelope.Contributor != "apikey" || envelope.Intent != "apikeys.list" { + t.Errorf("wrong forwarded request: %+v", envelope) + } + _ = json.NewEncoder(writer).Encode(map[string]any{"ok": true, "envelope": "v1", "kind": "query", "data": map[string]any{"keys": []string{"key-1"}}}) + })) + defer upstream.Close() + extension := New() + extension.clientMode = true + extension.config = Config{PortalURL: upstream.URL + "/authsome", ServiceAPIKey: "service-key"} + extension.SetLogger(forge.NewNoopLogger()) + registry := contract.NewRegistry() + remoteDispatcher := dispatcher.New(dispatcher.NoopMetricsEmitter{}) + if err := extension.registerRemoteContractContributor(remoteDispatcher, registry, contract.NewWardenRegistry()); err != nil { + t.Fatal(err) + } + for _, name := range []string{"auth", "apikey"} { + if !registry.IsRemote(name) { + t.Fatalf("%s is not registered remotely", name) + } + } + data, _, err := remoteDispatcher.Dispatch(context.Background(), contract.Request{Envelope: "v1", Kind: contract.KindQuery, Contributor: "apikey", Intent: "apikeys.list", IntentVersion: 1, Payload: json.RawMessage(`{}`)}, contract.Principal{}) + if err != nil { + t.Fatal(err) + } + if string(data) != `{"keys":["key-1"]}` { + t.Fatalf("unexpected response: %s", data) + } +} diff --git a/extension/extension.go b/extension/extension.go index 98c9ffd5..64ea0d48 100644 --- a/extension/extension.go +++ b/extension/extension.go @@ -949,7 +949,7 @@ func (e *Extension) registerRemoteContractContributor( // tries them, which is no worse than the pre-slice-(m) baseline. ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - m, err := contractremote.FetchManifest(ctx, remoteBaseURL, e.config.ServiceAPIKey, nil) + manifests, err := fetchContractCatalog(ctx, remoteBaseURL, e.config.ServiceAPIKey) if err != nil { // Error (not Warn) because this is the only signal that auth.* will // 404 at request time; operators need to see it in default log @@ -961,14 +961,18 @@ func (e *Extension) registerRemoteContractContributor( ) return nil //nolint:nilerr // non-fatal; surfaced via log } - if err := contractloader.Validate(m, wreg); err != nil { - return fmt.Errorf("authsome: validate remote manifest: %w", err) + for _, manifest := range manifests { + if err := contractloader.Validate(manifest, wreg); err != nil { + return fmt.Errorf("authsome: validate remote manifest %s: %w", manifest.Contributor.Name, err) + } } - if err := reg.RegisterRemote(m, dashcontract.RemoteEndpoint{ - BaseURL: remoteBaseURL, - APIKey: e.config.ServiceAPIKey, - }); err != nil { - return fmt.Errorf("authsome: register remote contributor: %w", err) + for _, manifest := range manifests { + if err := reg.RegisterRemote(manifest, dashcontract.RemoteEndpoint{ + BaseURL: remoteBaseURL, + APIKey: e.config.ServiceAPIKey, + }); err != nil { + return fmt.Errorf("authsome: register remote contributor %s: %w", manifest.Contributor.Name, err) + } } // Install the forwarding dispatcher. The dispatcher reads endpoints // from the registry per request, so this single install routes ALL @@ -977,7 +981,7 @@ func (e *Extension) registerRemoteContractContributor( // repeated calls just replace the field with an equivalent value. disp.SetRemoteDispatcher(contractremote.NewForwardingDispatcher(reg)) e.Logger().Info("authsome: registered upstream as remote contract contributor", - log.String("contributor", m.Contributor.Name), + log.Int("contributors", len(manifests)), log.String("remote_base", remoteBaseURL), ) return nil @@ -1001,18 +1005,7 @@ func (e *Extension) registerContractServer(router forge.Router) error { e.contractReg = dashcontract.NewRegistry() e.contractWreg = dashcontract.NewWardenRegistry() e.contractDisp = dispatcher.New(dispatcher.NoopMetricsEmitter{}) - deps := authcontract.Deps{ - Engine: e.engine, - SocialBasePath: e.config.Dashboard.SocialBasePath, - Brand: e.config.Dashboard.Brand, - BrandLogoURL: e.config.Dashboard.BrandLogoURL, - SignupURL: e.config.Dashboard.SignupURL, - SignupLabel: e.config.Dashboard.SignupLabel, - TermsURL: e.config.Dashboard.TermsURL, - PrivacyURL: e.config.Dashboard.PrivacyURL, - RequiredRoles: append([]string(nil), e.config.Dashboard.RequiredRoles...), - } - if err := authcontract.Register(e.contractDisp, e.contractReg, e.contractWreg, deps); err != nil { + if err := e.RegisterContractContributor(e.contractDisp, e.contractReg, e.contractWreg); err != nil { return fmt.Errorf("authsome: stand up contract server: %w", err) } srv := contractserver.New(e.contractReg, e.contractWreg, e.contractDisp, dashcontract.NoopAuditEmitter{}) diff --git a/go.mod b/go.mod index 044e4eea..1583d566 100644 --- a/go.mod +++ b/go.mod @@ -14,8 +14,8 @@ require ( github.com/stretchr/testify v1.12.1 github.com/testcontainers/testcontainers-go/modules/postgres v0.44.0 github.com/xraph/chronicle v1.6.2 - github.com/xraph/forge v1.11.0 - github.com/xraph/forge/extensions/auth v1.11.0 + github.com/xraph/forge v1.11.1 + github.com/xraph/forge/extensions/auth v1.11.1 github.com/xraph/forgeui v1.4.1 github.com/xraph/grove v1.6.3 github.com/xraph/grove/drivers/mongodriver v1.6.3 diff --git a/go.sum b/go.sum index eea526e9..3b4fcee1 100644 --- a/go.sum +++ b/go.sum @@ -430,10 +430,10 @@ github.com/xraph/confy v1.0.3 h1:mY0kIo9+dU+XcGGN7cm4IfOhvk4SDj5TrnccpGi21fs= github.com/xraph/confy v1.0.3/go.mod h1:DcYB+N0yKpr3t1yhhkdIIiCfch30Kh5wW6GRntUMOGk= github.com/xraph/dispatch v1.6.2 h1:pfyKiPuS1tIlhao+FyBlg36p6J+a5rTSMBTCE6gZlvM= github.com/xraph/dispatch v1.6.2/go.mod h1:K2lGkHo2U4EdVWVSdm22NtaCARSVc2y1OG5WYJsqp3E= -github.com/xraph/forge v1.11.0 h1:WUswdf2elUIQGvFjKKBDFjTmNRHa2kNzeEEKQ/92vrA= -github.com/xraph/forge v1.11.0/go.mod h1:aJ1aDFO7k5Js5oJwwHr0DJy09fyeMHbPzK5cF0tG1GU= -github.com/xraph/forge/extensions/auth v1.11.0 h1:EHh4l7SIDVnoL8xGn0QKYFzclYNqj/wl1KCEhm02I4A= -github.com/xraph/forge/extensions/auth v1.11.0/go.mod h1:iE2c9DLjt38rCfc5HfsT8LWtLAqDLFF5U33Jh+meuIQ= +github.com/xraph/forge v1.11.1 h1:dLjV0i3qZHEnaziI2t81ibco8hdL+GrX/bLFWB90mBE= +github.com/xraph/forge v1.11.1/go.mod h1:aJ1aDFO7k5Js5oJwwHr0DJy09fyeMHbPzK5cF0tG1GU= +github.com/xraph/forge/extensions/auth v1.11.1 h1:MrzgTwcTOCNPejHAZ4oYQGZJu0CbxN8IPPFoxxJkSeA= +github.com/xraph/forge/extensions/auth v1.11.1/go.mod h1:iE2c9DLjt38rCfc5HfsT8LWtLAqDLFF5U33Jh+meuIQ= github.com/xraph/forgeui v1.4.1 h1:LHK1t/sZ+9zL+MNUZralO9/rc0f5UCa19dpbWTuRMNg= github.com/xraph/forgeui v1.4.1/go.mod h1:rH/+wb1tt2pXSHotWAvoP+Lt846xlIjuwPDSpS5K5mw= github.com/xraph/go-utils v1.3.0 h1:bMZxwH9Xi/6UqhHSotKHB3a+/FnqnlUnEKf+aHZaw5c= diff --git a/plugins/anomaly/contract/handlers_test.go b/plugins/anomaly/contract/handlers_test.go index 7772860b..af93edb3 100644 --- a/plugins/anomaly/contract/handlers_test.go +++ b/plugins/anomaly/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/apikey/contract/handlers_test.go b/plugins/apikey/contract/handlers_test.go index 0b48e627..6ba7d084 100644 --- a/plugins/apikey/contract/handlers_test.go +++ b/plugins/apikey/contract/handlers_test.go @@ -21,9 +21,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 4 { t.Errorf("intents = %d, want 4 (list/detail/create/revoke)", got) } - if got := len(m.Graph); got != 3 { - t.Errorf("graph routes = %d, want 3 (/apikeys + /apikeys/:id + /apikeys/create)", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/consent/contract/handlers_test.go b/plugins/consent/contract/handlers_test.go index cd86e396..60681c98 100644 --- a/plugins/consent/contract/handlers_test.go +++ b/plugins/consent/contract/handlers_test.go @@ -21,9 +21,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 4 { t.Errorf("intents = %d, want 4 (list/userConsents/grant/revoke)", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1 (/compliance/consent)", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/deviceverify/contract/handlers_test.go b/plugins/deviceverify/contract/handlers_test.go index 28dbb81e..cd34f827 100644 --- a/plugins/deviceverify/contract/handlers_test.go +++ b/plugins/deviceverify/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/email/contract/handlers_test.go b/plugins/email/contract/handlers_test.go index e9037dbf..f4f69445 100644 --- a/plugins/email/contract/handlers_test.go +++ b/plugins/email/contract/handlers_test.go @@ -21,13 +21,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - // One graph route: /auth/email deep-link page. - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1 (/auth/email)", got) - } - if got := len(m.Extends); got != 0 { - t.Errorf("extends = %d, want 0 (auto-discovered via settings.tabs)", got) - } } // TestManifest_Validates ensures every intent referenced by the graph diff --git a/plugins/geofence/contract/handlers_test.go b/plugins/geofence/contract/handlers_test.go index 3dc706d1..027ddc0e 100644 --- a/plugins/geofence/contract/handlers_test.go +++ b/plugins/geofence/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/geoip/contract/handlers_test.go b/plugins/geoip/contract/handlers_test.go index c0d7276e..20f12ae7 100644 --- a/plugins/geoip/contract/handlers_test.go +++ b/plugins/geoip/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/impossibletravel/contract/handlers_test.go b/plugins/impossibletravel/contract/handlers_test.go index 16a30539..9cd57ef4 100644 --- a/plugins/impossibletravel/contract/handlers_test.go +++ b/plugins/impossibletravel/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/ipreputation/contract/handlers_test.go b/plugins/ipreputation/contract/handlers_test.go index 69f3bec9..a50743ce 100644 --- a/plugins/ipreputation/contract/handlers_test.go +++ b/plugins/ipreputation/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/magiclink/contract/handlers_test.go b/plugins/magiclink/contract/handlers_test.go index 5490c57f..ddf93ab9 100644 --- a/plugins/magiclink/contract/handlers_test.go +++ b/plugins/magiclink/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/mfa/contract/handlers_test.go b/plugins/mfa/contract/handlers_test.go index ebffb206..85240227 100644 --- a/plugins/mfa/contract/handlers_test.go +++ b/plugins/mfa/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/notification/contract.go b/plugins/notification/contract.go index 79f1894f..4ffafeba 100644 --- a/plugins/notification/contract.go +++ b/plugins/notification/contract.go @@ -4,8 +4,10 @@ package notification import ( "fmt" + "sort" authsome "github.com/xraph/authsome" + "github.com/xraph/authsome/bridge" "github.com/xraph/authsome/plugin" notifcontract "github.com/xraph/authsome/plugins/notification/contract" @@ -25,5 +27,19 @@ func (p *Plugin) RegisterContract( if !ok { return fmt.Errorf("notification: contract registration requires *authsome.Engine, got %T", engine) } - return notifcontract.Register(d, reg, wreg, notifcontract.Deps{Engine: eng}) + return notifcontract.Register(d, reg, wreg, notifcontract.Deps{ + Engine: eng, + Manager: func() bridge.HeraldTemplateManager { return p.templates }, + Sender: func() bridge.Herald { return p.herald }, + Mappings: func() []notifcontract.MappingSummary { + out := make([]notifcontract.MappingSummary, 0, len(p.mappings)) + for action, mapping := range p.mappings { + if mapping != nil { + out = append(out, notifcontract.MappingSummary{Action: action, Template: mapping.Template, Channels: append([]string(nil), mapping.Channels...), Enabled: mapping.Enabled}) + } + } + sort.Slice(out, func(i, j int) bool { return out[i].Action < out[j].Action }) + return out + }, + }) } diff --git a/plugins/notification/contract/contract.go b/plugins/notification/contract/contract.go index e8922b7c..f311ca85 100644 --- a/plugins/notification/contract/contract.go +++ b/plugins/notification/contract/contract.go @@ -8,6 +8,7 @@ import ( "fmt" authsome "github.com/xraph/authsome" + "github.com/xraph/authsome/bridge" "github.com/xraph/forge/extensions/dashboard/contract" "github.com/xraph/forge/extensions/dashboard/contract/dispatcher" @@ -18,7 +19,10 @@ import ( var manifestYAML []byte type Deps struct { - Engine *authsome.Engine + Engine *authsome.Engine + Manager func() bridge.HeraldTemplateManager + Sender func() bridge.Herald + Mappings func() []MappingSummary } func Register( @@ -40,6 +44,42 @@ func Register( if err := reg.Register(m); err != nil { return fmt.Errorf("notification/contract: register manifest: %w", err) } - _ = d + const c = "notification" + if err := dispatcher.RegisterQuery(d, c, "notification.templates.list", 1, templatesList(deps)); err != nil { + return err + } + if err := dispatcher.RegisterQuery(d, c, "notification.templates.detail", 1, templateDetail(deps)); err != nil { + return err + } + if err := dispatcher.RegisterQuery(d, c, "notification.templates.preview", 1, templatePreview(deps)); err != nil { + return err + } + if err := dispatcher.RegisterCommand(d, c, "notification.templates.create", 1, templateCreate(deps)); err != nil { + return err + } + if err := dispatcher.RegisterCommand(d, c, "notification.templates.update", 1, templateUpdate(deps)); err != nil { + return err + } + if err := dispatcher.RegisterCommand(d, c, "notification.templates.delete", 1, templateDelete(deps)); err != nil { + return err + } + if err := dispatcher.RegisterCommand(d, c, "notification.versions.create", 1, versionCreate(deps)); err != nil { + return err + } + if err := dispatcher.RegisterCommand(d, c, "notification.versions.update", 1, versionUpdate(deps)); err != nil { + return err + } + if err := dispatcher.RegisterCommand(d, c, "notification.versions.delete", 1, versionDelete(deps)); err != nil { + return err + } + if err := dispatcher.RegisterCommand(d, c, "notification.send", 1, sendNotification(deps)); err != nil { + return err + } + if err := dispatcher.RegisterCommand(d, c, "notification.templates.resetDefaults", 1, resetDefaultTemplates(deps)); err != nil { + return err + } + if err := dispatcher.RegisterQuery(d, c, "notification.mappings.list", 1, mappingsList(deps)); err != nil { + return err + } return nil } diff --git a/plugins/notification/contract/handlers.go b/plugins/notification/contract/handlers.go new file mode 100644 index 00000000..26fc1133 --- /dev/null +++ b/plugins/notification/contract/handlers.go @@ -0,0 +1,328 @@ +package contract + +import ( + "context" + "errors" + "strings" + + "github.com/xraph/authsome/bridge" + authcontract "github.com/xraph/authsome/extension/contract" + "github.com/xraph/forge/extensions/dashboard/contract" +) + +type templateIDInput struct { + ID string `json:"id"` +} +type previewInput struct { + ID string `json:"id"` + Locale string `json:"locale"` + Data map[string]any `json:"data"` +} +type createTemplateInput struct { + Name string `json:"name"` + Slug string `json:"slug"` + Channel string `json:"channel"` + Category string `json:"category"` + Locale string `json:"locale"` + Subject string `json:"subject"` + Title string `json:"title"` + HTML string `json:"html"` + Text string `json:"text"` +} +type updateTemplateInput struct { + ID string `json:"id"` + Name string `json:"name"` + Category string `json:"category"` + Enabled bool `json:"enabled"` +} +type versionInput struct { + TemplateID string `json:"templateId"` + ID string `json:"id"` + Locale string `json:"locale"` + Subject string `json:"subject"` + Title string `json:"title"` + HTML string `json:"html"` + Text string `json:"text"` + Active bool `json:"active"` +} +type sendInput struct { + ID string `json:"id"` + Recipient string `json:"recipient"` + Locale string `json:"locale"` + Data map[string]any `json:"data"` + Test bool `json:"test"` +} +type ack struct { + OK bool `json:"ok"` + ID string `json:"id,omitempty"` +} +type templatesResponse struct { + Templates []*bridge.HeraldTemplate `json:"templates"` +} +type MappingSummary struct { + Action string `json:"action"` + Template string `json:"template"` + Channels []string `json:"channels"` + Enabled bool `json:"enabled"` +} +type mappingsResponse struct { + Mappings []MappingSummary `json:"mappings"` +} + +func mappingsList(deps Deps) func(context.Context, struct{}, contract.Principal) (mappingsResponse, error) { + return func(_ context.Context, _ struct{}, _ contract.Principal) (mappingsResponse, error) { + if deps.Mappings == nil { + return mappingsResponse{Mappings: []MappingSummary{}}, nil + } + return mappingsResponse{Mappings: deps.Mappings()}, nil + } +} + +func invalid(message string) error { + return &contract.Error{Code: contract.CodeBadRequest, Message: message} +} +func missing() error { + return &contract.Error{Code: contract.CodeNotFound, Message: "template not found"} +} +func unavailable() error { + return &contract.Error{Code: contract.CodeUnavailable, Message: "notification service not configured"} +} +func failure(err error) error { + if err == nil { + return nil + } + var known *contract.Error + if errors.As(err, &known) { + return known + } + return &contract.Error{Code: contract.CodeInternal, Message: err.Error()} +} +func manager(deps Deps) (bridge.HeraldTemplateManager, error) { + if deps.Manager == nil { + return nil, unavailable() + } + m := deps.Manager() + if m == nil { + return nil, unavailable() + } + return m, nil +} +func appID(deps Deps, p contract.Principal) string { + return authcontract.AppIDFromPrincipal(p, deps.Engine).String() +} +func scopedTemplate(ctx context.Context, deps Deps, p contract.Principal, id string) (*bridge.HeraldTemplate, error) { + if strings.TrimSpace(id) == "" { + return nil, invalid("id is required") + } + m, err := manager(deps) + if err != nil { + return nil, err + } + t, err := m.GetTemplate(ctx, id) + if err != nil || t == nil { + return nil, missing() + } + if t.AppID != "" && t.AppID != appID(deps, p) { + return nil, missing() + } + return t, nil +} +func templatesList(deps Deps) func(context.Context, struct{}, contract.Principal) (templatesResponse, error) { + return func(ctx context.Context, _ struct{}, p contract.Principal) (templatesResponse, error) { + m, err := manager(deps) + if err != nil { + return templatesResponse{}, err + } + list, err := m.ListTemplates(ctx, appID(deps, p)) + if err != nil { + return templatesResponse{}, failure(err) + } + if list == nil { + list = []*bridge.HeraldTemplate{} + } + return templatesResponse{Templates: list}, nil + } +} +func templateDetail(deps Deps) func(context.Context, templateIDInput, contract.Principal) (*bridge.HeraldTemplate, error) { + return func(ctx context.Context, in templateIDInput, p contract.Principal) (*bridge.HeraldTemplate, error) { + return scopedTemplate(ctx, deps, p, in.ID) + } +} +func templatePreview(deps Deps) func(context.Context, previewInput, contract.Principal) (*bridge.HeraldRenderedContent, error) { + return func(ctx context.Context, in previewInput, p contract.Principal) (*bridge.HeraldRenderedContent, error) { + if _, err := scopedTemplate(ctx, deps, p, in.ID); err != nil { + return nil, err + } + m, _ := manager(deps) + out, err := m.RenderTemplate(ctx, in.ID, in.Locale, in.Data) + return out, failure(err) + } +} +func templateCreate(deps Deps) func(context.Context, createTemplateInput, contract.Principal) (ack, error) { + return func(ctx context.Context, in createTemplateInput, p contract.Principal) (ack, error) { + if strings.TrimSpace(in.Name) == "" || strings.TrimSpace(in.Slug) == "" { + return ack{}, invalid("name and slug are required") + } + switch in.Channel { + case "email", "sms", "inapp", "push": + default: + return ack{}, invalid("invalid channel") + } + m, err := manager(deps) + if err != nil { + return ack{}, err + } + locale := strings.TrimSpace(in.Locale) + if locale == "" { + locale = "en" + } + category := strings.TrimSpace(in.Category) + if category == "" { + category = bridge.HeraldCategoryTransactional + } + t := &bridge.HeraldTemplate{AppID: appID(deps, p), Name: strings.TrimSpace(in.Name), Slug: strings.TrimSpace(in.Slug), Channel: in.Channel, Category: category, Enabled: true, + Versions: []bridge.HeraldTemplateVersion{{Locale: locale, Subject: in.Subject, Title: in.Title, HTML: in.HTML, Text: in.Text, Active: true}}} + if err := m.CreateTemplate(ctx, t); err != nil { + return ack{}, failure(err) + } + return ack{OK: true, ID: t.ID}, nil + } +} +func templateUpdate(deps Deps) func(context.Context, updateTemplateInput, contract.Principal) (ack, error) { + return func(ctx context.Context, in updateTemplateInput, p contract.Principal) (ack, error) { + t, err := scopedTemplate(ctx, deps, p, in.ID) + if err != nil { + return ack{}, err + } + if strings.TrimSpace(in.Name) == "" { + return ack{}, invalid("name is required") + } + t.Name = strings.TrimSpace(in.Name) + t.Category = strings.TrimSpace(in.Category) + t.Enabled = in.Enabled + m, _ := manager(deps) + if err := m.UpdateTemplate(ctx, t); err != nil { + return ack{}, failure(err) + } + return ack{OK: true, ID: t.ID}, nil + } +} +func templateDelete(deps Deps) func(context.Context, templateIDInput, contract.Principal) (ack, error) { + return func(ctx context.Context, in templateIDInput, p contract.Principal) (ack, error) { + t, err := scopedTemplate(ctx, deps, p, in.ID) + if err != nil { + return ack{}, err + } + if t.IsSystem { + return ack{}, invalid("system templates cannot be deleted") + } + m, _ := manager(deps) + if err := m.DeleteTemplate(ctx, in.ID); err != nil { + return ack{}, failure(err) + } + return ack{OK: true, ID: in.ID}, nil + } +} +func versionFor(t *bridge.HeraldTemplate, id string) bool { + for _, version := range t.Versions { + if version.ID == id { + return true + } + } + return false +} +func versionCreate(deps Deps) func(context.Context, versionInput, contract.Principal) (ack, error) { + return func(ctx context.Context, in versionInput, p contract.Principal) (ack, error) { + if _, err := scopedTemplate(ctx, deps, p, in.TemplateID); err != nil { + return ack{}, err + } + if strings.TrimSpace(in.Locale) == "" { + return ack{}, invalid("locale is required") + } + m, _ := manager(deps) + v := &bridge.HeraldTemplateVersion{TemplateID: in.TemplateID, Locale: strings.TrimSpace(in.Locale), Subject: in.Subject, Title: in.Title, HTML: in.HTML, Text: in.Text, Active: true} + if err := m.CreateVersion(ctx, v); err != nil { + return ack{}, failure(err) + } + return ack{OK: true, ID: v.ID}, nil + } +} +func versionUpdate(deps Deps) func(context.Context, versionInput, contract.Principal) (ack, error) { + return func(ctx context.Context, in versionInput, p contract.Principal) (ack, error) { + t, err := scopedTemplate(ctx, deps, p, in.TemplateID) + if err != nil { + return ack{}, err + } + if !versionFor(t, in.ID) { + return ack{}, missing() + } + if strings.TrimSpace(in.Locale) == "" { + return ack{}, invalid("locale is required") + } + m, _ := manager(deps) + v := &bridge.HeraldTemplateVersion{ID: in.ID, TemplateID: in.TemplateID, Locale: in.Locale, Subject: in.Subject, Title: in.Title, HTML: in.HTML, Text: in.Text, Active: in.Active} + if err := m.UpdateVersion(ctx, v); err != nil { + return ack{}, failure(err) + } + return ack{OK: true, ID: in.ID}, nil + } +} +func versionDelete(deps Deps) func(context.Context, versionInput, contract.Principal) (ack, error) { + return func(ctx context.Context, in versionInput, p contract.Principal) (ack, error) { + t, err := scopedTemplate(ctx, deps, p, in.TemplateID) + if err != nil { + return ack{}, err + } + if !versionFor(t, in.ID) { + return ack{}, missing() + } + m, _ := manager(deps) + if err := m.DeleteVersion(ctx, in.ID); err != nil { + return ack{}, failure(err) + } + return ack{OK: true, ID: in.ID}, nil + } +} +func sendNotification(deps Deps) func(context.Context, sendInput, contract.Principal) (ack, error) { + return func(ctx context.Context, in sendInput, p contract.Principal) (ack, error) { + t, err := scopedTemplate(ctx, deps, p, in.ID) + if err != nil { + return ack{}, err + } + if !t.Enabled { + return ack{}, invalid("template is disabled") + } + recipient := strings.TrimSpace(in.Recipient) + if recipient == "" { + return ack{}, invalid("recipient is required") + } + req := &bridge.HeraldSendRequest{AppID: appID(deps, p), Channel: t.Channel, Template: t.Slug, Locale: in.Locale, To: []string{recipient}, Data: in.Data} + if in.Test { + m, _ := manager(deps) + req.Metadata = map[string]string{"test": "true"} + if err := m.TestSend(ctx, req); err != nil { + return ack{}, failure(err) + } + } else { + if deps.Sender == nil || deps.Sender() == nil { + return ack{}, unavailable() + } + if err := deps.Sender().Send(ctx, req); err != nil { + return ack{}, failure(err) + } + } + return ack{OK: true, ID: t.ID}, nil + } +} +func resetDefaultTemplates(deps Deps) func(context.Context, struct{}, contract.Principal) (ack, error) { + return func(ctx context.Context, _ struct{}, p contract.Principal) (ack, error) { + m, err := manager(deps) + if err != nil { + return ack{}, err + } + if err := m.ResetDefaultTemplates(ctx, appID(deps, p)); err != nil { + return ack{}, failure(err) + } + return ack{OK: true}, nil + } +} diff --git a/plugins/notification/contract/handlers_test.go b/plugins/notification/contract/handlers_test.go index c80e4a68..ff3c1726 100644 --- a/plugins/notification/contract/handlers_test.go +++ b/plugins/notification/contract/handlers_test.go @@ -2,12 +2,69 @@ package contract import ( "bytes" + "context" + "errors" "testing" + authsome "github.com/xraph/authsome" + "github.com/xraph/authsome/bridge" + "github.com/xraph/authsome/id" "github.com/xraph/forge/extensions/dashboard/contract" "github.com/xraph/forge/extensions/dashboard/contract/loader" ) +type managerStub struct { + bridge.HeraldTemplateManager + template *bridge.HeraldTemplate + testRequest *bridge.HeraldSendRequest +} + +func (m *managerStub) GetTemplate(_ context.Context, _ string) (*bridge.HeraldTemplate, error) { + return m.template, nil +} +func (m *managerStub) TestSend(_ context.Context, req *bridge.HeraldSendRequest) error { + m.testRequest = req + return nil +} + +type senderStub struct { + bridge.Herald + request *bridge.HeraldSendRequest +} + +func (s *senderStub) Send(_ context.Context, req *bridge.HeraldSendRequest) error { + s.request = req + return nil +} + +func TestSendNotificationUsesScopedTemplateAndTestTransport(t *testing.T) { + app := id.NewAppID().String() + other := id.NewAppID().String() + m := &managerStub{template: &bridge.HeraldTemplate{ID: "template-1", AppID: app, Slug: "welcome", Channel: "email", Enabled: true}} + s := &senderStub{} + deps := Deps{Engine: &authsome.Engine{}, Manager: func() bridge.HeraldTemplateManager { return m }, Sender: func() bridge.Herald { return s }} + p := contract.Principal{Claims: map[string]any{"app_id": other}} + _, err := sendNotification(deps)(context.Background(), sendInput{ID: "template-1", Recipient: "user@example.com"}, p) + var ce *contract.Error + if !errors.As(err, &ce) || ce.Code != contract.CodeNotFound { + t.Fatalf("cross-app template should be hidden, got %v", err) + } + if s.request != nil { + t.Fatal("cross-app notification was sent") + } + p.Claims["app_id"] = app + result, err := sendNotification(deps)(context.Background(), sendInput{ID: "template-1", Recipient: "user@example.com", Locale: "en", Data: map[string]any{"name": "Ada"}, Test: true}, p) + if err != nil || !result.OK { + t.Fatalf("test send: result=%+v err=%v", result, err) + } + if m.testRequest == nil || m.testRequest.AppID != app || m.testRequest.Template != "welcome" || m.testRequest.Channel != "email" || len(m.testRequest.To) != 1 || m.testRequest.To[0] != "user@example.com" { + t.Fatalf("unexpected test send request: %+v", m.testRequest) + } + if s.request != nil { + t.Fatal("test send reached live sender") + } +} + func TestManifest_Loads(t *testing.T) { m, err := loader.Load(bytes.NewReader(manifestYAML), "notification/contract/manifest.yaml") if err != nil { @@ -16,11 +73,8 @@ func TestManifest_Loads(t *testing.T) { if m.Contributor.Name != "notification" { t.Errorf("contributor name = %q, want notification", m.Contributor.Name) } - if got := len(m.Intents); got != 0 { - t.Errorf("intents = %d, want 0", got) - } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) + if got := len(m.Intents); got != 12 { + t.Errorf("intents = %d, want 12", got) } } diff --git a/plugins/notification/contract/manifest.yaml b/plugins/notification/contract/manifest.yaml index 8598ca17..4a280cb3 100644 --- a/plugins/notification/contract/manifest.yaml +++ b/plugins/notification/contract/manifest.yaml @@ -4,9 +4,21 @@ contributor: envelope: supports: [v1] preferred: v1 - capabilities: [auth.read] + capabilities: [auth.read, auth.write] -intents: [] +intents: + - { name: notification.templates.list, kind: query, version: 1, capability: read } + - { name: notification.templates.detail, kind: query, version: 1, capability: read } + - { name: notification.templates.preview, kind: query, version: 1, capability: read } + - { name: notification.templates.create, kind: command, version: 1, capability: write, invalidates: [notification.templates.list] } + - { name: notification.templates.update, kind: command, version: 1, capability: write, invalidates: [notification.templates.list, notification.templates.detail] } + - { name: notification.templates.delete, kind: command, version: 1, capability: write, invalidates: [notification.templates.list, notification.templates.detail] } + - { name: notification.versions.create, kind: command, version: 1, capability: write, invalidates: [notification.templates.detail, notification.templates.preview] } + - { name: notification.versions.update, kind: command, version: 1, capability: write, invalidates: [notification.templates.detail, notification.templates.preview] } + - { name: notification.versions.delete, kind: command, version: 1, capability: write, invalidates: [notification.templates.detail, notification.templates.preview] } + - { name: notification.send, kind: command, version: 1, capability: write } + - { name: notification.templates.resetDefaults, kind: command, version: 1, capability: write, invalidates: [notification.templates.list] } + - { name: notification.mappings.list, kind: query, version: 1, capability: read } graph: - route: /notifications diff --git a/plugins/oauth2provider/contract/handlers_test.go b/plugins/oauth2provider/contract/handlers_test.go index 8df934d0..20551bfe 100644 --- a/plugins/oauth2provider/contract/handlers_test.go +++ b/plugins/oauth2provider/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/organization/contract/contract.go b/plugins/organization/contract/contract.go index 8897e6cd..881d6642 100644 --- a/plugins/organization/contract/contract.go +++ b/plugins/organization/contract/contract.go @@ -37,7 +37,10 @@ type OrgService interface { UpdateOrganization(ctx context.Context, o *organization.Organization) error DeleteOrganization(ctx context.Context, orgID id.OrgID) error ListMembers(ctx context.Context, orgID id.OrgID) ([]*organization.Member, error) + AddMember(ctx context.Context, m *organization.Member) error RemoveMember(ctx context.Context, memberID id.MemberID) error + ListInvitations(ctx context.Context, orgID id.OrgID) ([]*organization.Invitation, error) + CreateInvitation(ctx context.Context, inv *organization.Invitation) error } // Deps carries the typed plugin handle alongside the engine so the @@ -91,8 +94,17 @@ func Register( if err := dispatcher.RegisterQuery(d, c, "orgs.members", 1, orgsMembersListHandler(deps)); err != nil { return fmt.Errorf("organization/contract: register orgs.members: %w", err) } + if err := dispatcher.RegisterCommand(d, c, "orgs.addMember", 1, orgsAddMemberHandler(deps)); err != nil { + return fmt.Errorf("organization/contract: register orgs.addMember: %w", err) + } if err := dispatcher.RegisterCommand(d, c, "orgs.removeMember", 1, orgsRemoveMemberHandler(deps)); err != nil { return fmt.Errorf("organization/contract: register orgs.removeMember: %w", err) } + if err := dispatcher.RegisterQuery(d, c, "orgs.invitations", 1, orgsInvitationsHandler(deps)); err != nil { + return fmt.Errorf("organization/contract: register orgs.invitations: %w", err) + } + if err := dispatcher.RegisterCommand(d, c, "orgs.createInvitation", 1, orgsCreateInvitationHandler(deps)); err != nil { + return fmt.Errorf("organization/contract: register orgs.createInvitation: %w", err) + } return nil } diff --git a/plugins/organization/contract/handlers.go b/plugins/organization/contract/handlers.go index cb480e43..5da5067b 100644 --- a/plugins/organization/contract/handlers.go +++ b/plugins/organization/contract/handlers.go @@ -8,11 +8,15 @@ package contract import ( "context" "errors" + "net/mail" "strings" "time" + "github.com/xraph/authsome/account" "github.com/xraph/authsome/id" "github.com/xraph/authsome/organization" + "github.com/xraph/authsome/store" + "github.com/xraph/authsome/user" "github.com/xraph/forge/extensions/dashboard/contract" @@ -53,6 +57,24 @@ type MembersListResponse struct { Members []MemberSummary `json:"members"` } +type InvitationSummary struct { + ID string `json:"id"` + Email string `json:"email"` + Role string `json:"role"` + Status string `json:"status"` + CreatedAt string `json:"createdAt"` + ExpiresAt string `json:"expiresAt"` +} + +type InvitationsListResponse struct { + Invitations []InvitationSummary `json:"invitations"` +} + +type CreateInvitationResponse struct { + InvitationSummary + Token string `json:"token"` +} + type GetOrgInput struct { ID string `json:"id"` } @@ -77,6 +99,23 @@ type ListMembersInput struct { OrgID string `json:"orgId"` } +type AddMemberInput struct { + OrgID string `json:"orgId"` + UserID string `json:"userId"` + Email string `json:"email,omitempty"` + Role string `json:"role,omitempty"` +} + +type ListInvitationsInput struct { + OrgID string `json:"orgId"` +} + +type CreateInvitationInput struct { + OrgID string `json:"orgId"` + Email string `json:"email"` + Role string `json:"role,omitempty"` +} + type RemoveMemberInput struct { ID string `json:"id"` } @@ -216,6 +255,126 @@ func orgsMembersListHandler(deps Deps) func(ctx context.Context, in ListMembersI } } +func orgsAddMemberHandler(deps Deps) func(context.Context, AddMemberInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in AddMemberInput, p contract.Principal) (ackResponse, error) { + if deps.Engine == nil || deps.Plugin == nil { + return ackResponse{}, unavailable() + } + org, err := scopedOrg(ctx, deps, in.OrgID, p) + if err != nil { + return ackResponse{}, err + } + role, err := memberRole(in.Role) + if err != nil { + return ackResponse{}, err + } + var u *user.User + if rawID := strings.TrimSpace(in.UserID); rawID != "" { + uid, parseErr := id.ParseUserID(rawID) + if parseErr != nil { + return ackResponse{}, badReq("valid userId is required") + } + u, err = deps.Engine.GetUser(ctx, uid) + } else if email := strings.ToLower(strings.TrimSpace(in.Email)); email != "" { + address, parseErr := mail.ParseAddress(email) + if parseErr != nil || address.Address != email { + return ackResponse{}, badReq("valid email is required") + } + u, err = deps.Engine.GetUserByEmail(ctx, org.AppID, email) + } else { + return ackResponse{}, badReq("userId or email is required") + } + if err != nil { + if errors.Is(err, store.ErrNotFound) { + return ackResponse{}, &contract.Error{Code: contract.CodeNotFound, Message: "user not found in this app"} + } + return ackResponse{}, mapErr(err) + } + if u == nil || u.AppID != org.AppID { + return ackResponse{}, &contract.Error{Code: contract.CodeNotFound, Message: "user not found in this app"} + } + uid := u.ID + members, err := deps.Plugin.ListMembers(ctx, org.ID) + if err != nil { + return ackResponse{}, mapErr(err) + } + for _, member := range members { + if member.UserID == uid { + return ackResponse{}, &contract.Error{Code: contract.CodeConflict, Message: "user is already a member"} + } + } + member := &organization.Member{ID: id.NewMemberID(), OrgID: org.ID, UserID: uid, Role: role} + if err := deps.Plugin.AddMember(ctx, member); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: member.ID.String()}, nil + } +} + +func orgsInvitationsHandler(deps Deps) func(context.Context, ListInvitationsInput, contract.Principal) (InvitationsListResponse, error) { + return func(ctx context.Context, in ListInvitationsInput, p contract.Principal) (InvitationsListResponse, error) { + if deps.Engine == nil || deps.Plugin == nil { + return InvitationsListResponse{}, unavailable() + } + org, err := scopedOrg(ctx, deps, in.OrgID, p) + if err != nil { + return InvitationsListResponse{}, err + } + list, err := deps.Plugin.ListInvitations(ctx, org.ID) + if err != nil { + return InvitationsListResponse{}, mapErr(err) + } + out := InvitationsListResponse{Invitations: make([]InvitationSummary, 0, len(list))} + for _, inv := range list { + out.Invitations = append(out.Invitations, projectInvitation(inv)) + } + return out, nil + } +} + +func orgsCreateInvitationHandler(deps Deps) func(context.Context, CreateInvitationInput, contract.Principal) (CreateInvitationResponse, error) { + return func(ctx context.Context, in CreateInvitationInput, p contract.Principal) (CreateInvitationResponse, error) { + if deps.Engine == nil || deps.Plugin == nil { + return CreateInvitationResponse{}, unavailable() + } + org, err := scopedOrg(ctx, deps, in.OrgID, p) + if err != nil { + return CreateInvitationResponse{}, err + } + inviterID, err := principalUserID(p) + if err != nil { + return CreateInvitationResponse{}, err + } + email := strings.ToLower(strings.TrimSpace(in.Email)) + address, err := mail.ParseAddress(email) + if err != nil || address.Address != email { + return CreateInvitationResponse{}, badReq("valid email is required") + } + role, err := memberRole(in.Role) + if err != nil { + return CreateInvitationResponse{}, err + } + list, err := deps.Plugin.ListInvitations(ctx, org.ID) + if err != nil { + return CreateInvitationResponse{}, mapErr(err) + } + for _, existing := range list { + if strings.EqualFold(existing.Email, email) && existing.Status == organization.InvitationPending && (existing.ExpiresAt.IsZero() || time.Now().Before(existing.ExpiresAt)) { + return CreateInvitationResponse{}, &contract.Error{Code: contract.CodeConflict, Message: "a pending invitation already exists for this email"} + } + } + token, err := account.GenerateVerificationToken() + if err != nil { + return CreateInvitationResponse{}, mapErr(err) + } + inv := &organization.Invitation{ID: id.NewInvitationID(), OrgID: org.ID, Email: email, Role: role, InviterID: inviterID, Status: organization.InvitationPending, Token: token, ExpiresAt: time.Now().Add(72 * time.Hour)} + if err := deps.Plugin.CreateInvitation(ctx, inv); err != nil { + return CreateInvitationResponse{}, mapErr(err) + } + return CreateInvitationResponse{InvitationSummary: projectInvitation(inv), Token: token}, nil + } +} + func orgsRemoveMemberHandler(deps Deps) func(ctx context.Context, in RemoveMemberInput, _ contract.Principal) (ackResponse, error) { return func(ctx context.Context, in RemoveMemberInput, _ contract.Principal) (ackResponse, error) { if deps.Engine == nil || deps.Plugin == nil { @@ -259,6 +418,44 @@ func projectOrgDetail(o *organization.Organization) OrgDetail { } } +func projectInvitation(inv *organization.Invitation) InvitationSummary { + return InvitationSummary{ + ID: inv.ID.String(), Email: inv.Email, Role: string(inv.Role), Status: string(inv.Status), + CreatedAt: inv.CreatedAt.UTC().Format(time.RFC3339), ExpiresAt: inv.ExpiresAt.UTC().Format(time.RFC3339), + } +} + +func scopedOrg(ctx context.Context, deps Deps, rawID string, p contract.Principal) (*organization.Organization, error) { + oid, err := parseOrgID(rawID) + if err != nil { + return nil, err + } + org, err := deps.Plugin.GetOrganization(ctx, oid) + if err != nil { + if errors.Is(err, store.ErrNotFound) { + return nil, &contract.Error{Code: contract.CodeNotFound, Message: "organization not found in this app"} + } + return nil, mapErr(err) + } + if org == nil || org.AppID != authcontract.AppIDFromPrincipal(p, deps.Engine) { + return nil, &contract.Error{Code: contract.CodeNotFound, Message: "organization not found in this app"} + } + return org, nil +} + +func memberRole(raw string) (organization.MemberRole, error) { + if raw == "" { + return organization.RoleMember, nil + } + role := organization.MemberRole(raw) + switch role { + case organization.RoleMember, organization.RoleAdmin, organization.RoleOwner: + return role, nil + default: + return "", badReq("role must be member, admin, or owner") + } +} + func parseOrgID(s string) (id.OrgID, error) { if strings.TrimSpace(s) == "" { return id.OrgID{}, badReq("id is required") diff --git a/plugins/organization/contract/handlers_test.go b/plugins/organization/contract/handlers_test.go index 5c8047a1..78502cde 100644 --- a/plugins/organization/contract/handlers_test.go +++ b/plugins/organization/contract/handlers_test.go @@ -3,13 +3,51 @@ package contract import ( "bytes" "context" + "encoding/json" "errors" "testing" + authsome "github.com/xraph/authsome" + "github.com/xraph/authsome/id" + "github.com/xraph/authsome/organization" + "github.com/xraph/authsome/store/memory" + "github.com/xraph/authsome/user" + dashauth "github.com/xraph/forge/extensions/dashboard/auth" "github.com/xraph/forge/extensions/dashboard/contract" "github.com/xraph/forge/extensions/dashboard/contract/loader" + "github.com/xraph/warden" + wardenmem "github.com/xraph/warden/store/memory" ) +type invitationService struct { + OrgService + org *organization.Organization + invitations []*organization.Invitation + members []*organization.Member +} + +func (s *invitationService) GetOrganization(_ context.Context, _ id.OrgID) (*organization.Organization, error) { + return s.org, nil +} + +func (s *invitationService) ListInvitations(_ context.Context, _ id.OrgID) ([]*organization.Invitation, error) { + return s.invitations, nil +} + +func (s *invitationService) CreateInvitation(_ context.Context, inv *organization.Invitation) error { + s.invitations = append(s.invitations, inv) + return nil +} + +func (s *invitationService) ListMembers(_ context.Context, _ id.OrgID) ([]*organization.Member, error) { + return s.members, nil +} + +func (s *invitationService) AddMember(_ context.Context, member *organization.Member) error { + s.members = append(s.members, member) + return nil +} + func TestManifest_Loads(t *testing.T) { m, err := loader.Load(bytes.NewReader(manifestYAML), "organization/contract/manifest.yaml") if err != nil { @@ -18,11 +56,8 @@ func TestManifest_Loads(t *testing.T) { if m.Contributor.Name != "organization" { t.Errorf("contributor name = %q, want organization", m.Contributor.Name) } - if got := len(m.Intents); got != 7 { - t.Errorf("intents = %d, want 7 (list/detail/create/update/delete/members/removeMember)", got) - } - if got := len(m.Graph); got != 3 { - t.Errorf("graph routes = %d, want 3 (/organizations + /:id + /create)", got) + if got := len(m.Intents); got != 10 { + t.Errorf("intents = %d, want 10", got) } } @@ -44,3 +79,66 @@ func TestOrgsListHandler_UnavailableWhenEngineNil(t *testing.T) { t.Errorf("expected CodeUnavailable, got %v", err) } } + +func TestInvitationHandlers_OneTimeTokenAndAppScope(t *testing.T) { + w, err := warden.NewEngine(warden.WithStore(wardenmem.New())) + if err != nil { + t.Fatal(err) + } + mem := memory.New() + engine, err := authsome.NewEngine(authsome.WithStore(mem), authsome.WithWarden(w)) + if err != nil { + t.Fatal(err) + } + appID := id.NewAppID() + org := &organization.Organization{ID: id.NewOrgID(), AppID: appID} + service := &invitationService{org: org} + deps := Deps{Engine: engine, Plugin: service} + p := contract.Principal{User: &dashauth.UserInfo{Subject: id.NewUserID().String()}, Claims: map[string]any{"app_id": appID.String()}} + existing := &user.User{ID: id.NewUserID(), AppID: appID, Email: "member@example.com"} + if err := mem.CreateUserWithPrimaryEmail(context.Background(), existing, user.NewPrimaryEmail(existing, "admin")); err != nil { + t.Fatal(err) + } + added, err := orgsAddMemberHandler(deps)(context.Background(), AddMemberInput{OrgID: org.ID.String(), Email: " Member@Example.com "}, p) + if err != nil || !added.OK || len(service.members) != 1 || service.members[0].UserID != existing.ID { + t.Fatalf("add existing user by email: response=%+v members=%+v err=%v", added, service.members, err) + } + + created, err := orgsCreateInvitationHandler(deps)(context.Background(), CreateInvitationInput{OrgID: org.ID.String(), Email: " Person@Example.com ", Role: "admin"}, p) + if err != nil { + t.Fatal(err) + } + if created.Token == "" || created.Email != "person@example.com" || created.Role != "admin" { + t.Fatalf("unexpected invitation: %+v", created) + } + if created.ExpiresAt == "" { + t.Fatal("invitation should expire") + } + listed, err := orgsInvitationsHandler(deps)(context.Background(), ListInvitationsInput{OrgID: org.ID.String()}, p) + if err != nil { + t.Fatal(err) + } + encoded, err := json.Marshal(listed) + if err != nil { + t.Fatal(err) + } + if len(listed.Invitations) != 1 || string(encoded) == "" || containsToken(encoded, created.Token) { + t.Fatalf("list leaked token or lost invitation: %s", encoded) + } + _, err = orgsCreateInvitationHandler(deps)(context.Background(), CreateInvitationInput{OrgID: org.ID.String(), Email: "person@example.com"}, p) + var ce *contract.Error + if !errors.As(err, &ce) || ce.Code != contract.CodeConflict { + t.Fatalf("duplicate invitation: %v", err) + } + + otherApp := id.NewAppID() + p.Claims["app_id"] = otherApp.String() + _, err = orgsInvitationsHandler(deps)(context.Background(), ListInvitationsInput{OrgID: org.ID.String()}, p) + if !errors.As(err, &ce) || ce.Code != contract.CodeNotFound { + t.Fatalf("cross-app list: %v", err) + } +} + +func containsToken(data []byte, token string) bool { + return bytes.Contains(data, []byte(token)) +} diff --git a/plugins/organization/contract/manifest.yaml b/plugins/organization/contract/manifest.yaml index 86d8e1cd..a237d695 100644 --- a/plugins/organization/contract/manifest.yaml +++ b/plugins/organization/contract/manifest.yaml @@ -17,7 +17,10 @@ intents: - { name: orgs.update, kind: command, version: 1, capability: write, invalidates: [orgs.list, orgs.detail] } - { name: orgs.delete, kind: command, version: 1, capability: write, invalidates: [orgs.list] } - { name: orgs.members, kind: query, version: 1, capability: read } + - { name: orgs.addMember, kind: command, version: 1, capability: write, invalidates: [orgs.members] } - { name: orgs.removeMember, kind: command, version: 1, capability: write, invalidates: [orgs.members] } + - { name: orgs.invitations, kind: query, version: 1, capability: read } + - { name: orgs.createInvitation, kind: command, version: 1, capability: write, invalidates: [orgs.invitations] } queries: orgList: diff --git a/plugins/passkey/contract/handlers_test.go b/plugins/passkey/contract/handlers_test.go index 8fdc81b9..1a80608e 100644 --- a/plugins/passkey/contract/handlers_test.go +++ b/plugins/passkey/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/password/contract/handlers_test.go b/plugins/password/contract/handlers_test.go index 6ac63235..340b7331 100644 --- a/plugins/password/contract/handlers_test.go +++ b/plugins/password/contract/handlers_test.go @@ -26,16 +26,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 1 { t.Errorf("intents = %d, want 1 (policy)", got) } - // One graph route: /auth/password (deep-link page rendering - // the settings.panel for the password namespace). - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1 (/auth/password)", got) - } - // No extends — the global /settings page auto-discovers via - // settings.tabs + settings.namespaces. - if got := len(m.Extends); got != 0 { - t.Errorf("extends = %d, want 0 (no manual extension; auto-discovered)", got) - } } // TestManifest_Validates ensures every intent referenced by the graph diff --git a/plugins/phone/contract/handlers_test.go b/plugins/phone/contract/handlers_test.go index b03b82bc..1a9db294 100644 --- a/plugins/phone/contract/handlers_test.go +++ b/plugins/phone/contract/handlers_test.go @@ -19,12 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1 (/auth/phone)", got) - } - if got := len(m.Extends); got != 0 { - t.Errorf("extends = %d, want 0 (auto-discovered)", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/riskengine/contract/handlers_test.go b/plugins/riskengine/contract/handlers_test.go index 7cd54e58..03dea71f 100644 --- a/plugins/riskengine/contract/handlers_test.go +++ b/plugins/riskengine/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/scim/contract/handlers_test.go b/plugins/scim/contract/handlers_test.go index 26f0f4ce..79f4f7e8 100644 --- a/plugins/scim/contract/handlers_test.go +++ b/plugins/scim/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/social/contract/handlers_test.go b/plugins/social/contract/handlers_test.go index dd9accc2..ba4c7513 100644 --- a/plugins/social/contract/handlers_test.go +++ b/plugins/social/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/sso/contract/handlers_test.go b/plugins/sso/contract/handlers_test.go index 27bb0352..c9ea3df0 100644 --- a/plugins/sso/contract/handlers_test.go +++ b/plugins/sso/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/subscription/contract.go b/plugins/subscription/contract.go index 23178fcb..04812d87 100644 --- a/plugins/subscription/contract.go +++ b/plugins/subscription/contract.go @@ -3,6 +3,7 @@ package subscription import ( + "context" "fmt" authsome "github.com/xraph/authsome" @@ -29,5 +30,15 @@ func (p *Plugin) RegisterContract( if svc == nil { return fmt.Errorf("subscription: Service not initialised") } - return subcontract.Register(d, reg, wreg, subcontract.Deps{Engine: eng, Service: svc}) + return subcontract.Register(d, reg, wreg, subcontract.Deps{Engine: eng, Service: svc, Usage: func(ctx context.Context, tenantID, appID string) ([]subcontract.UsageSummary, error) { + rows, err := svc.GetUsageSummary(ctx, tenantID, appID) + if err != nil { + return nil, err + } + out := make([]subcontract.UsageSummary, 0, len(rows)) + for _, row := range rows { + out = append(out, subcontract.UsageSummary{FeatureKey: row.FeatureKey, FeatureName: row.FeatureName, FeatureType: row.FeatureType, Used: row.Used, Limit: row.Limit, Remaining: row.Remaining, Period: row.Period}) + } + return out, nil + }}) } diff --git a/plugins/subscription/contract/contract.go b/plugins/subscription/contract/contract.go index 4f15c6cd..2566d8c0 100644 --- a/plugins/subscription/contract/contract.go +++ b/plugins/subscription/contract/contract.go @@ -10,7 +10,10 @@ import ( _ "embed" "fmt" + "github.com/xraph/ledger/coupon" + "github.com/xraph/ledger/feature" ledgerid "github.com/xraph/ledger/id" + "github.com/xraph/ledger/invoice" "github.com/xraph/ledger/plan" "github.com/xraph/ledger/subscription" @@ -32,14 +35,39 @@ var manifestYAML []byte type SubscriptionService interface { ListPlans(ctx context.Context, appID string) ([]*plan.Plan, error) GetPlan(ctx context.Context, planID ledgerid.PlanID) (*plan.Plan, error) + CreatePlan(ctx context.Context, p *plan.Plan) error + UpdatePlan(ctx context.Context, p *plan.Plan) error ArchivePlan(ctx context.Context, planID ledgerid.PlanID) error ActivatePlan(ctx context.Context, planID ledgerid.PlanID) error ListSubscriptions(ctx context.Context, tenantID, appID string, opts subscription.ListOpts) ([]*subscription.Subscription, error) + GetSubscription(ctx context.Context, subID ledgerid.SubscriptionID) (*subscription.Subscription, error) + GetActiveSubscription(ctx context.Context, tenantID, appID string) (*subscription.Subscription, error) + Subscribe(ctx context.Context, tenantID string, planID ledgerid.PlanID, appID string) (*subscription.Subscription, error) + ChangePlan(ctx context.Context, subID ledgerid.SubscriptionID, newPlanID ledgerid.PlanID) error + PauseSubscription(ctx context.Context, subID ledgerid.SubscriptionID) error + ResumeSubscription(ctx context.Context, subID ledgerid.SubscriptionID) error + CancelSubscription(ctx context.Context, subID ledgerid.SubscriptionID, immediately bool) error + ListAllInvoices(ctx context.Context, appID string) ([]*invoice.Invoice, error) + ListInvoices(ctx context.Context, tenantID, appID string) ([]*invoice.Invoice, error) + GetInvoice(ctx context.Context, invID ledgerid.InvoiceID) (*invoice.Invoice, error) + GenerateInvoice(ctx context.Context, subID ledgerid.SubscriptionID) (*invoice.Invoice, error) + MarkInvoicePaid(ctx context.Context, invID ledgerid.InvoiceID, paymentRef string) error + MarkInvoiceVoided(ctx context.Context, invID ledgerid.InvoiceID, reason string) error + ListCoupons(ctx context.Context, appID string) ([]*coupon.Coupon, error) + GetCoupon(ctx context.Context, code, appID string) (*coupon.Coupon, error) + CreateCoupon(ctx context.Context, c *coupon.Coupon) error + DeleteCoupon(ctx context.Context, couponID ledgerid.CouponID) error + ListCatalogFeatures(ctx context.Context, appID string) ([]*feature.Feature, error) + GetCatalogFeature(ctx context.Context, featureID ledgerid.FeatureID) (*feature.Feature, error) + CreateCatalogFeature(ctx context.Context, f *feature.Feature) error + UpdateCatalogFeature(ctx context.Context, f *feature.Feature) error + ArchiveCatalogFeature(ctx context.Context, featureID ledgerid.FeatureID) error } type Deps struct { Engine *authsome.Engine Service SubscriptionService + Usage func(context.Context, string, string) ([]UsageSummary, error) } func Register( @@ -81,5 +109,57 @@ func Register( if err := dispatcher.RegisterQuery(d, c, "subscriptions.list", 1, subscriptionsListHandler(deps)); err != nil { return fmt.Errorf("subscription/contract: register subscriptions.list: %w", err) } + registrations := []func() error{ + func() error { return dispatcher.RegisterCommand(d, c, "plans.create", 1, plansCreateHandler(deps)) }, + func() error { return dispatcher.RegisterCommand(d, c, "plans.update", 1, plansUpdateHandler(deps)) }, + func() error { + return dispatcher.RegisterQuery(d, c, "subscriptions.all", 1, subscriptionsAllHandler(deps)) + }, + func() error { + return dispatcher.RegisterQuery(d, c, "subscriptions.detail", 1, subscriptionsDetailHandler(deps)) + }, + func() error { + return dispatcher.RegisterCommand(d, c, "subscriptions.create", 1, subscriptionsCreateHandler(deps)) + }, + func() error { + return dispatcher.RegisterCommand(d, c, "subscriptions.change", 1, subscriptionsChangeHandler(deps)) + }, + func() error { + return dispatcher.RegisterCommand(d, c, "subscriptions.pause", 1, subscriptionsPauseHandler(deps)) + }, + func() error { + return dispatcher.RegisterCommand(d, c, "subscriptions.resume", 1, subscriptionsResumeHandler(deps)) + }, + func() error { + return dispatcher.RegisterCommand(d, c, "subscriptions.cancel", 1, subscriptionsCancelHandler(deps)) + }, + func() error { return dispatcher.RegisterQuery(d, c, "invoices.list", 1, invoicesListHandler(deps)) }, + func() error { return dispatcher.RegisterQuery(d, c, "invoices.detail", 1, invoicesDetailHandler(deps)) }, + func() error { + return dispatcher.RegisterCommand(d, c, "invoices.generate", 1, invoicesGenerateHandler(deps)) + }, + func() error { + return dispatcher.RegisterCommand(d, c, "invoices.markPaid", 1, invoicesMarkPaidHandler(deps)) + }, + func() error { return dispatcher.RegisterCommand(d, c, "invoices.void", 1, invoicesVoidHandler(deps)) }, + func() error { return dispatcher.RegisterQuery(d, c, "coupons.list", 1, couponsListHandler(deps)) }, + func() error { return dispatcher.RegisterCommand(d, c, "coupons.create", 1, couponsCreateHandler(deps)) }, + func() error { return dispatcher.RegisterCommand(d, c, "coupons.delete", 1, couponsDeleteHandler(deps)) }, + func() error { return dispatcher.RegisterQuery(d, c, "features.list", 1, featuresListHandler(deps)) }, + func() error { + return dispatcher.RegisterCommand(d, c, "features.create", 1, featuresCreateHandler(deps)) + }, + func() error { + return dispatcher.RegisterCommand(d, c, "features.update", 1, featuresUpdateHandler(deps)) + }, + func() error { + return dispatcher.RegisterCommand(d, c, "features.archive", 1, featuresArchiveHandler(deps)) + }, + } + for _, register := range registrations { + if err := register(); err != nil { + return fmt.Errorf("subscription/contract: register billing intent: %w", err) + } + } return nil } diff --git a/plugins/subscription/contract/handlers.go b/plugins/subscription/contract/handlers.go index 77da752c..fcc4efaf 100644 --- a/plugins/subscription/contract/handlers.go +++ b/plugins/subscription/contract/handlers.go @@ -24,13 +24,15 @@ import ( // ──────────────────────────────────────────────────────────────────── type PlanSummary struct { - ID string `json:"id"` - Name string `json:"name"` - Slug string `json:"slug"` - Description string `json:"description,omitempty"` - Currency string `json:"currency,omitempty"` - Status string `json:"status"` - TrialDays int `json:"trialDays,omitempty"` + ID string `json:"id"` + Name string `json:"name"` + Slug string `json:"slug"` + Description string `json:"description,omitempty"` + Currency string `json:"currency,omitempty"` + Status string `json:"status"` + TrialDays int `json:"trialDays,omitempty"` + BaseAmount int64 `json:"baseAmount"` + BillingPeriod string `json:"billingPeriod,omitempty"` } type SubscriptionSummary struct { @@ -40,19 +42,32 @@ type SubscriptionSummary struct { Status string `json:"status"` CurrentPeriodStart string `json:"currentPeriodStart,omitempty"` CurrentPeriodEnd string `json:"currentPeriodEnd,omitempty"` + CancelAt string `json:"cancelAt,omitempty"` } type PlanDetail struct { PlanSummary - Features []PlanFeature `json:"features,omitempty"` + Features []PlanFeature `json:"features,omitempty"` + Tiers []PriceTierSummary `json:"tiers,omitempty"` } type PlanFeature struct { - Key string `json:"key"` - Name string `json:"name"` - Type string `json:"type"` - Limit int64 `json:"limit"` - Period string `json:"period"` + ID string `json:"id,omitempty"` + Key string `json:"key"` + Name string `json:"name"` + Type string `json:"type"` + Limit int64 `json:"limit"` + Period string `json:"period"` + SoftLimit bool `json:"softLimit"` + CatalogID string `json:"catalogId,omitempty"` +} + +type PriceTierSummary struct { + FeatureKey string `json:"featureKey"` + Type string `json:"type"` + UpTo int64 `json:"upTo"` + UnitAmount int64 `json:"unitAmount"` + FlatAmount int64 `json:"flatAmount"` } type PlansListResponse struct { @@ -105,61 +120,64 @@ func plansListHandler(deps Deps) func(ctx context.Context, _ struct{}, p contrac } } -func plansDetailHandler(deps Deps) func(ctx context.Context, in GetPlanInput, _ contract.Principal) (PlanDetail, error) { - return func(ctx context.Context, in GetPlanInput, _ contract.Principal) (PlanDetail, error) { +func plansDetailHandler(deps Deps) func(ctx context.Context, in GetPlanInput, p contract.Principal) (PlanDetail, error) { + return func(ctx context.Context, in GetPlanInput, p contract.Principal) (PlanDetail, error) { if deps.Engine == nil || deps.Service == nil { return PlanDetail{}, unavailable() } - pid, err := parsePlanID(in.ID) - if err != nil { - return PlanDetail{}, err - } - pl, err := deps.Service.GetPlan(ctx, pid) + pl, err := scopedPlan(ctx, deps, p, in.ID) if err != nil { return PlanDetail{}, mapErr(err) } d := PlanDetail{PlanSummary: projectPlan(pl)} for _, f := range pl.Features { d.Features = append(d.Features, PlanFeature{ - Key: f.Key, Name: f.Name, - Type: string(f.Type), - Limit: f.Limit, - Period: string(f.Period), + ID: f.ID.String(), Key: f.Key, Name: f.Name, + Type: string(f.Type), + Limit: f.Limit, + Period: string(f.Period), + SoftLimit: f.SoftLimit, + CatalogID: f.CatalogID.String(), }) } + if pl.Pricing != nil { + for _, tier := range pl.Pricing.Tiers { + d.Tiers = append(d.Tiers, PriceTierSummary{FeatureKey: tier.FeatureKey, Type: string(tier.Type), UpTo: tier.UpTo, UnitAmount: tier.UnitAmount.Amount, FlatAmount: tier.FlatAmount.Amount}) + } + } return d, nil } } -func plansArchiveHandler(deps Deps) func(ctx context.Context, in ArchivePlanInput, _ contract.Principal) (ackResponse, error) { - return func(ctx context.Context, in ArchivePlanInput, _ contract.Principal) (ackResponse, error) { +func plansArchiveHandler(deps Deps) func(ctx context.Context, in ArchivePlanInput, p contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in ArchivePlanInput, p contract.Principal) (ackResponse, error) { if deps.Engine == nil || deps.Service == nil { return ackResponse{}, unavailable() } - pid, err := parsePlanID(in.ID) + pl, err := scopedPlan(ctx, deps, p, in.ID) if err != nil { return ackResponse{}, err } - if err := deps.Service.ArchivePlan(ctx, pid); err != nil { + if err := deps.Service.ArchivePlan(ctx, pl.ID); err != nil { return ackResponse{}, mapErr(err) } - return ackResponse{OK: true, ID: pid.String()}, nil + return ackResponse{OK: true, ID: pl.ID.String()}, nil } } -func plansActivateHandler(deps Deps) func(ctx context.Context, in ActivatePlanInput, _ contract.Principal) (ackResponse, error) { - return func(ctx context.Context, in ActivatePlanInput, _ contract.Principal) (ackResponse, error) { +func plansActivateHandler(deps Deps) func(ctx context.Context, in ActivatePlanInput, p contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in ActivatePlanInput, p contract.Principal) (ackResponse, error) { if deps.Engine == nil || deps.Service == nil { return ackResponse{}, unavailable() } - pid, err := parsePlanID(in.ID) + pl, err := scopedPlan(ctx, deps, p, in.ID) if err != nil { return ackResponse{}, err } - if err := deps.Service.ActivatePlan(ctx, pid); err != nil { + if err := deps.Service.ActivatePlan(ctx, pl.ID); err != nil { return ackResponse{}, mapErr(err) } - return ackResponse{OK: true, ID: pid.String()}, nil + return ackResponse{OK: true, ID: pl.ID.String()}, nil } } @@ -178,7 +196,9 @@ func subscriptionsListHandler(deps Deps) func(ctx context.Context, in Subscripti } out := SubscriptionsListResponse{Subscriptions: make([]SubscriptionSummary, 0, len(list))} for _, s := range list { - out.Subscriptions = append(out.Subscriptions, projectSubscription(s)) + if s.AppID == authcontract.AppIDFromPrincipal(p, deps.Engine).String() { + out.Subscriptions = append(out.Subscriptions, projectSubscription(s)) + } } return out, nil } @@ -192,23 +212,32 @@ func projectPlan(p *plan.Plan) PlanSummary { if p == nil { return PlanSummary{} } - return PlanSummary{ + out := PlanSummary{ ID: p.ID.String(), Name: p.Name, Slug: p.Slug, Description: p.Description, Currency: p.Currency, Status: string(p.Status), TrialDays: p.TrialDays, } + if p.Pricing != nil { + out.BaseAmount = p.Pricing.BaseAmount.Amount + out.BillingPeriod = string(p.Pricing.BillingPeriod) + } + return out } func projectSubscription(s *subscription.Subscription) SubscriptionSummary { if s == nil { return SubscriptionSummary{} } - return SubscriptionSummary{ + out := SubscriptionSummary{ ID: s.ID.String(), TenantID: s.TenantID, PlanID: s.PlanID.String(), Status: string(s.Status), CurrentPeriodStart: s.CurrentPeriodStart.UTC().Format(time.RFC3339), CurrentPeriodEnd: s.CurrentPeriodEnd.UTC().Format(time.RFC3339), } + if s.CancelAt != nil { + out.CancelAt = s.CancelAt.UTC().Format(time.RFC3339) + } + return out } func parsePlanID(s string) (ledgerid.PlanID, error) { diff --git a/plugins/subscription/contract/handlers_billing.go b/plugins/subscription/contract/handlers_billing.go new file mode 100644 index 00000000..f10eea66 --- /dev/null +++ b/plugins/subscription/contract/handlers_billing.go @@ -0,0 +1,693 @@ +package contract + +import ( + "context" + "strings" + "time" + + "github.com/xraph/ledger/coupon" + "github.com/xraph/ledger/feature" + ledgerid "github.com/xraph/ledger/id" + "github.com/xraph/ledger/invoice" + "github.com/xraph/ledger/plan" + "github.com/xraph/ledger/subscription" + "github.com/xraph/ledger/types" + + authcontract "github.com/xraph/authsome/extension/contract" + "github.com/xraph/forge/extensions/dashboard/contract" +) + +type planWriteInput struct { + ID string `json:"id"` + Name string `json:"name"` + Slug string `json:"slug"` + Description string `json:"description"` + Currency string `json:"currency"` + TrialDays int `json:"trialDays"` + BaseAmount int64 `json:"baseAmount"` + BillingPeriod string `json:"billingPeriod"` + Features []PlanFeature `json:"features"` + Tiers []PriceTierSummary `json:"tiers"` +} +type subscriptionWriteInput struct { + ID string `json:"id"` + TenantID string `json:"tenantId"` + PlanID string `json:"planId"` + Immediately bool `json:"immediately"` +} +type invoiceWriteInput struct { + ID string `json:"id"` + SubscriptionID string `json:"subscriptionId"` + PaymentRef string `json:"paymentRef"` + Reason string `json:"reason"` +} +type couponWriteInput struct { + ID string `json:"id"` + Code string `json:"code"` + Name string `json:"name"` + Type string `json:"type"` + Amount int64 `json:"amount"` + Percentage int `json:"percentage"` + Currency string `json:"currency"` + MaxRedemptions int `json:"maxRedemptions"` + ValidFrom string `json:"validFrom"` + ValidUntil string `json:"validUntil"` +} +type invoicesResponse struct { + Invoices []invoiceSummary `json:"invoices"` +} +type detailInput struct { + ID string `json:"id"` +} +type UsageSummary struct { + FeatureKey string `json:"featureKey"` + FeatureName string `json:"featureName"` + FeatureType string `json:"featureType"` + Used int64 `json:"used"` + Limit int64 `json:"limit"` + Remaining int64 `json:"remaining"` + Period string `json:"period"` +} +type subscriptionDetail struct { + SubscriptionSummary + PlanName string `json:"planName"` + Usage []UsageSummary `json:"usage"` + Invoices []invoiceSummary `json:"invoices"` +} +type invoiceLineSummary struct { + Description string `json:"description"` + Type string `json:"type"` + FeatureKey string `json:"featureKey,omitempty"` + Quantity int64 `json:"quantity"` + UnitAmount int64 `json:"unitAmount"` + Amount int64 `json:"amount"` +} +type invoiceDetail struct { + invoiceSummary + Subtotal int64 `json:"subtotal"` + TaxAmount int64 `json:"taxAmount"` + DiscountAmount int64 `json:"discountAmount"` + PeriodStart string `json:"periodStart"` + PeriodEnd string `json:"periodEnd"` + LineItems []invoiceLineSummary `json:"lineItems"` +} +type invoiceSummary struct { + ID string `json:"id"` + TenantID string `json:"tenantId"` + SubscriptionID string `json:"subscriptionId"` + Status string `json:"status"` + Currency string `json:"currency"` + Total int64 `json:"total"` + PaymentRef string `json:"paymentRef,omitempty"` +} + +func projectInvoice(item *invoice.Invoice) invoiceSummary { + return invoiceSummary{ID: item.ID.String(), TenantID: item.TenantID, SubscriptionID: item.SubscriptionID.String(), Status: string(item.Status), Currency: item.Currency, Total: item.Total.Amount, PaymentRef: item.PaymentRef} +} + +type couponsResponse struct { + Coupons []couponSummary `json:"coupons"` +} +type couponSummary struct { + ID string `json:"id"` + Code string `json:"code"` + Name string `json:"name"` + Type string `json:"type"` + Amount int64 `json:"amount"` + Percentage int `json:"percentage"` + Currency string `json:"currency"` + MaxRedemptions int `json:"maxRedemptions"` + TimesRedeemed int `json:"timesRedeemed"` + ValidFrom string `json:"validFrom,omitempty"` + ValidUntil string `json:"validUntil,omitempty"` +} + +func scopedApp(deps Deps, p contract.Principal) string { + return authcontract.AppIDFromPrincipal(p, deps.Engine).String() +} +func notFound(resource string) error { + return &contract.Error{Code: contract.CodeNotFound, Message: resource + " not found"} +} +func scopedPlan(ctx context.Context, deps Deps, p contract.Principal, raw string) (*plan.Plan, error) { + id, err := parsePlanID(raw) + if err != nil { + return nil, err + } + item, err := deps.Service.GetPlan(ctx, id) + if err != nil { + return nil, mapErr(err) + } + if item == nil || item.AppID != scopedApp(deps, p) { + return nil, notFound("plan") + } + return item, nil +} +func scopedSubscription(ctx context.Context, deps Deps, p contract.Principal, raw string) (*subscription.Subscription, error) { + id, err := ledgerid.ParseSubscriptionID(strings.TrimSpace(raw)) + if err != nil { + return nil, badReq("invalid subscription id") + } + item, err := deps.Service.GetSubscription(ctx, id) + if err != nil { + return nil, mapErr(err) + } + if item == nil || item.AppID != scopedApp(deps, p) { + return nil, notFound("subscription") + } + return item, nil +} +func scopedInvoice(ctx context.Context, deps Deps, p contract.Principal, raw string) (*invoice.Invoice, error) { + id, err := ledgerid.ParseInvoiceID(strings.TrimSpace(raw)) + if err != nil { + return nil, badReq("invalid invoice id") + } + item, err := deps.Service.GetInvoice(ctx, id) + if err != nil { + return nil, mapErr(err) + } + if item == nil || item.AppID != scopedApp(deps, p) { + return nil, notFound("invoice") + } + return item, nil +} +func pricing(amount int64, currency, period string, tiers []PriceTierSummary) (*plan.Pricing, error) { + if amount < 0 { + return nil, badReq("base amount cannot be negative") + } + if amount == 0 && len(tiers) == 0 { + return nil, nil + } + if period == "" { + period = string(plan.PeriodMonthly) + } + if period != string(plan.PeriodMonthly) && period != string(plan.PeriodYearly) { + return nil, badReq("invalid billing period") + } + out := &plan.Pricing{BaseAmount: types.Money{Amount: amount, Currency: currency}, BillingPeriod: plan.Period(period)} + for _, item := range tiers { + if strings.TrimSpace(item.FeatureKey) == "" || item.UpTo < 0 || item.UnitAmount < 0 || item.FlatAmount < 0 { + return nil, badReq("invalid price tier") + } + kind := plan.TierType(item.Type) + if kind != plan.TierGraduated && kind != plan.TierVolume && kind != plan.TierFlat { + return nil, badReq("invalid price tier type") + } + out.Tiers = append(out.Tiers, plan.PriceTier{FeatureKey: strings.TrimSpace(item.FeatureKey), Type: kind, UpTo: item.UpTo, UnitAmount: types.Money{Amount: item.UnitAmount, Currency: currency}, FlatAmount: types.Money{Amount: item.FlatAmount, Currency: currency}, Priority: len(out.Tiers)}) + } + return out, nil +} +func planFeatures(input []PlanFeature) ([]plan.Feature, error) { + out := make([]plan.Feature, 0, len(input)) + seen := map[string]bool{} + for _, item := range input { + key := strings.TrimSpace(item.Key) + if key == "" || seen[key] { + return nil, badReq("feature keys must be unique and nonempty") + } + seen[key] = true + kind := plan.FeatureType(item.Type) + if kind != plan.FeatureBoolean && kind != plan.FeatureMetered && kind != plan.FeatureSeat { + return nil, badReq("invalid feature type") + } + f := plan.Feature{ID: ledgerid.NewFeatureID(), Key: key, Name: strings.TrimSpace(item.Name), Type: kind, Limit: item.Limit, Period: plan.Period(item.Period), SoftLimit: item.SoftLimit} + if item.CatalogID != "" { + cid, err := ledgerid.ParseFeatureID(item.CatalogID) + if err != nil { + return nil, badReq("invalid catalog feature id") + } + f.CatalogID = cid + } + out = append(out, f) + } + return out, nil +} +func validateTierFeatures(price *plan.Pricing, features []plan.Feature) error { + if price == nil { + return nil + } + keys := make(map[string]bool, len(features)) + for _, item := range features { + keys[item.Key] = true + } + for _, tier := range price.Tiers { + if !keys[tier.FeatureKey] { + return badReq("price tier feature is not on the plan") + } + } + return nil +} +func plansCreateHandler(deps Deps) func(context.Context, planWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in planWriteInput, p contract.Principal) (ackResponse, error) { + if strings.TrimSpace(in.Name) == "" || strings.TrimSpace(in.Slug) == "" { + return ackResponse{}, badReq("name and slug are required") + } + if in.TrialDays < 0 { + return ackResponse{}, badReq("trial days cannot be negative") + } + currency := strings.ToLower(strings.TrimSpace(in.Currency)) + if currency == "" { + currency = "usd" + } + price, err := pricing(in.BaseAmount, currency, in.BillingPeriod, in.Tiers) + if err != nil { + return ackResponse{}, err + } + features, err := planFeatures(in.Features) + if err != nil { + return ackResponse{}, err + } + if err := validateTierFeatures(price, features); err != nil { + return ackResponse{}, err + } + item := &plan.Plan{Name: strings.TrimSpace(in.Name), Slug: strings.TrimSpace(in.Slug), Description: in.Description, Currency: currency, Status: plan.StatusDraft, TrialDays: in.TrialDays, AppID: scopedApp(deps, p), Pricing: price, Features: features} + if err := deps.Service.CreatePlan(ctx, item); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func plansUpdateHandler(deps Deps) func(context.Context, planWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in planWriteInput, p contract.Principal) (ackResponse, error) { + item, err := scopedPlan(ctx, deps, p, in.ID) + if err != nil { + return ackResponse{}, err + } + if strings.TrimSpace(in.Name) == "" { + return ackResponse{}, badReq("name is required") + } + if in.TrialDays < 0 { + return ackResponse{}, badReq("trial days cannot be negative") + } + price, err := pricing(in.BaseAmount, item.Currency, in.BillingPeriod, in.Tiers) + if err != nil { + return ackResponse{}, err + } + features, err := planFeatures(in.Features) + if err != nil { + return ackResponse{}, err + } + if err := validateTierFeatures(price, features); err != nil { + return ackResponse{}, err + } + oldIDs := map[string]ledgerid.FeatureID{} + for _, f := range item.Features { + oldIDs[f.Key] = f.ID + } + for i := range features { + if oldID, ok := oldIDs[features[i].Key]; ok { + features[i].ID = oldID + } + } + item.Name = strings.TrimSpace(in.Name) + item.Description = in.Description + item.TrialDays = in.TrialDays + item.Features = features + if price != nil && item.Pricing != nil { + price.ID = item.Pricing.ID + price.PlanID = item.Pricing.PlanID + } + item.Pricing = price + if err := deps.Service.UpdatePlan(ctx, item); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func subscriptionsAllHandler(deps Deps) func(context.Context, struct{}, contract.Principal) (SubscriptionsListResponse, error) { + return func(ctx context.Context, _ struct{}, p contract.Principal) (SubscriptionsListResponse, error) { + list, err := deps.Service.ListSubscriptions(ctx, "", scopedApp(deps, p), subscription.ListOpts{}) + if err != nil { + return SubscriptionsListResponse{}, mapErr(err) + } + out := SubscriptionsListResponse{Subscriptions: make([]SubscriptionSummary, 0, len(list))} + for _, item := range list { + if item.AppID == scopedApp(deps, p) { + out.Subscriptions = append(out.Subscriptions, projectSubscription(item)) + } + } + return out, nil + } +} +func subscriptionsDetailHandler(deps Deps) func(context.Context, detailInput, contract.Principal) (subscriptionDetail, error) { + return func(ctx context.Context, in detailInput, p contract.Principal) (subscriptionDetail, error) { + item, err := scopedSubscription(ctx, deps, p, in.ID) + if err != nil { + return subscriptionDetail{}, err + } + out := subscriptionDetail{SubscriptionSummary: projectSubscription(item), Usage: []UsageSummary{}, Invoices: []invoiceSummary{}} + pl, err := scopedPlan(ctx, deps, p, item.PlanID.String()) + if err == nil { + out.PlanName = pl.Name + } + if deps.Usage != nil { + active, activeErr := deps.Service.GetActiveSubscription(ctx, item.TenantID, scopedApp(deps, p)) + if activeErr == nil && active != nil && active.ID == item.ID { + if usage, err := deps.Usage(ctx, item.TenantID, scopedApp(deps, p)); err == nil { + out.Usage = usage + } + } + } + invoices, err := deps.Service.ListInvoices(ctx, item.TenantID, scopedApp(deps, p)) + if err != nil { + return subscriptionDetail{}, mapErr(err) + } + for _, inv := range invoices { + if inv.AppID == scopedApp(deps, p) && inv.SubscriptionID == item.ID { + out.Invoices = append(out.Invoices, projectInvoice(inv)) + } + } + return out, nil + } +} +func subscriptionsCreateHandler(deps Deps) func(context.Context, subscriptionWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in subscriptionWriteInput, p contract.Principal) (ackResponse, error) { + if strings.TrimSpace(in.TenantID) == "" { + return ackResponse{}, badReq("tenant id is required") + } + pl, err := scopedPlan(ctx, deps, p, in.PlanID) + if err != nil { + return ackResponse{}, err + } + if pl.Status != plan.StatusActive { + return ackResponse{}, badReq("plan must be active") + } + item, err := deps.Service.Subscribe(ctx, strings.TrimSpace(in.TenantID), pl.ID, scopedApp(deps, p)) + if err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func subscriptionsChangeHandler(deps Deps) func(context.Context, subscriptionWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in subscriptionWriteInput, p contract.Principal) (ackResponse, error) { + item, err := scopedSubscription(ctx, deps, p, in.ID) + if err != nil { + return ackResponse{}, err + } + pl, err := scopedPlan(ctx, deps, p, in.PlanID) + if err != nil { + return ackResponse{}, err + } + if pl.Status != plan.StatusActive { + return ackResponse{}, badReq("plan must be active") + } + if err := deps.Service.ChangePlan(ctx, item.ID, pl.ID); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func subscriptionAction(deps Deps, action func(context.Context, ledgerid.SubscriptionID) error) func(context.Context, subscriptionWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in subscriptionWriteInput, p contract.Principal) (ackResponse, error) { + item, err := scopedSubscription(ctx, deps, p, in.ID) + if err != nil { + return ackResponse{}, err + } + if err := action(ctx, item.ID); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func subscriptionsPauseHandler(deps Deps) func(context.Context, subscriptionWriteInput, contract.Principal) (ackResponse, error) { + return subscriptionAction(deps, deps.Service.PauseSubscription) +} +func subscriptionsResumeHandler(deps Deps) func(context.Context, subscriptionWriteInput, contract.Principal) (ackResponse, error) { + return subscriptionAction(deps, deps.Service.ResumeSubscription) +} +func subscriptionsCancelHandler(deps Deps) func(context.Context, subscriptionWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in subscriptionWriteInput, p contract.Principal) (ackResponse, error) { + item, err := scopedSubscription(ctx, deps, p, in.ID) + if err != nil { + return ackResponse{}, err + } + if err := deps.Service.CancelSubscription(ctx, item.ID, in.Immediately); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func invoicesListHandler(deps Deps) func(context.Context, struct{}, contract.Principal) (invoicesResponse, error) { + return func(ctx context.Context, _ struct{}, p contract.Principal) (invoicesResponse, error) { + list, err := deps.Service.ListAllInvoices(ctx, scopedApp(deps, p)) + if err != nil { + return invoicesResponse{}, mapErr(err) + } + out := invoicesResponse{Invoices: make([]invoiceSummary, 0, len(list))} + for _, item := range list { + if item.AppID == scopedApp(deps, p) { + out.Invoices = append(out.Invoices, invoiceSummary{ID: item.ID.String(), TenantID: item.TenantID, SubscriptionID: item.SubscriptionID.String(), Status: string(item.Status), Currency: item.Currency, Total: item.Total.Amount, PaymentRef: item.PaymentRef}) + } + } + return out, nil + } +} +func invoicesDetailHandler(deps Deps) func(context.Context, detailInput, contract.Principal) (invoiceDetail, error) { + return func(ctx context.Context, in detailInput, p contract.Principal) (invoiceDetail, error) { + item, err := scopedInvoice(ctx, deps, p, in.ID) + if err != nil { + return invoiceDetail{}, err + } + out := invoiceDetail{invoiceSummary: projectInvoice(item), Subtotal: item.Subtotal.Amount, TaxAmount: item.TaxAmount.Amount, DiscountAmount: item.DiscountAmount.Amount, PeriodStart: item.PeriodStart.UTC().Format("2006-01-02"), PeriodEnd: item.PeriodEnd.UTC().Format("2006-01-02"), LineItems: make([]invoiceLineSummary, 0, len(item.LineItems))} + for _, line := range item.LineItems { + out.LineItems = append(out.LineItems, invoiceLineSummary{Description: line.Description, Type: string(line.Type), FeatureKey: line.FeatureKey, Quantity: line.Quantity, UnitAmount: line.UnitAmount.Amount, Amount: line.Amount.Amount}) + } + return out, nil + } +} +func invoicesGenerateHandler(deps Deps) func(context.Context, invoiceWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in invoiceWriteInput, p contract.Principal) (ackResponse, error) { + item, err := scopedSubscription(ctx, deps, p, in.SubscriptionID) + if err != nil { + return ackResponse{}, err + } + generated, err := deps.Service.GenerateInvoice(ctx, item.ID) + if err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: generated.ID.String()}, nil + } +} +func invoicesMarkPaidHandler(deps Deps) func(context.Context, invoiceWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in invoiceWriteInput, p contract.Principal) (ackResponse, error) { + item, err := scopedInvoice(ctx, deps, p, in.ID) + if err != nil { + return ackResponse{}, err + } + if strings.TrimSpace(in.PaymentRef) == "" { + return ackResponse{}, badReq("payment reference is required") + } + if err := deps.Service.MarkInvoicePaid(ctx, item.ID, strings.TrimSpace(in.PaymentRef)); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func invoicesVoidHandler(deps Deps) func(context.Context, invoiceWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in invoiceWriteInput, p contract.Principal) (ackResponse, error) { + item, err := scopedInvoice(ctx, deps, p, in.ID) + if err != nil { + return ackResponse{}, err + } + if err := deps.Service.MarkInvoiceVoided(ctx, item.ID, strings.TrimSpace(in.Reason)); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func couponsListHandler(deps Deps) func(context.Context, struct{}, contract.Principal) (couponsResponse, error) { + return func(ctx context.Context, _ struct{}, p contract.Principal) (couponsResponse, error) { + list, err := deps.Service.ListCoupons(ctx, scopedApp(deps, p)) + if err != nil { + return couponsResponse{}, mapErr(err) + } + out := couponsResponse{Coupons: make([]couponSummary, 0, len(list))} + for _, item := range list { + if item.AppID == scopedApp(deps, p) { + row := couponSummary{ID: item.ID.String(), Code: item.Code, Name: item.Name, Type: string(item.Type), Amount: item.Amount.Amount, Percentage: item.Percentage, Currency: item.Currency, MaxRedemptions: item.MaxRedemptions, TimesRedeemed: item.TimesRedeemed} + if item.ValidFrom != nil { + row.ValidFrom = item.ValidFrom.Format(time.DateOnly) + } + if item.ValidUntil != nil { + row.ValidUntil = item.ValidUntil.Format(time.DateOnly) + } + out.Coupons = append(out.Coupons, row) + } + } + return out, nil + } +} +func couponsCreateHandler(deps Deps) func(context.Context, couponWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in couponWriteInput, p contract.Principal) (ackResponse, error) { + if strings.TrimSpace(in.Code) == "" || strings.TrimSpace(in.Name) == "" { + return ackResponse{}, badReq("code and name are required") + } + kind := coupon.CouponType(in.Type) + if kind != coupon.CouponTypePercentage && kind != coupon.CouponTypeAmount { + return ackResponse{}, badReq("invalid coupon type") + } + if kind == coupon.CouponTypePercentage && (in.Percentage <= 0 || in.Percentage > 100) { + return ackResponse{}, badReq("percentage must be between 1 and 100") + } + if kind == coupon.CouponTypeAmount && in.Amount <= 0 { + return ackResponse{}, badReq("amount must be positive") + } + currency := strings.ToLower(strings.TrimSpace(in.Currency)) + if currency == "" { + currency = "usd" + } + item := &coupon.Coupon{Code: strings.TrimSpace(in.Code), Name: strings.TrimSpace(in.Name), Type: kind, Amount: types.Money{Amount: in.Amount, Currency: currency}, Percentage: in.Percentage, Currency: currency, MaxRedemptions: in.MaxRedemptions, AppID: scopedApp(deps, p)} + if in.MaxRedemptions < 0 { + return ackResponse{}, badReq("max redemptions cannot be negative") + } + if in.ValidFrom != "" { + date, err := time.Parse(time.DateOnly, in.ValidFrom) + if err != nil { + return ackResponse{}, badReq("invalid start date") + } + item.ValidFrom = &date + } + if in.ValidUntil != "" { + date, err := time.Parse(time.DateOnly, in.ValidUntil) + if err != nil { + return ackResponse{}, badReq("invalid expiry date") + } + item.ValidUntil = &date + } + if item.ValidFrom != nil && item.ValidUntil != nil && item.ValidUntil.Before(*item.ValidFrom) { + return ackResponse{}, badReq("expiry must follow start date") + } + if err := deps.Service.CreateCoupon(ctx, item); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func couponsDeleteHandler(deps Deps) func(context.Context, couponWriteInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in couponWriteInput, p contract.Principal) (ackResponse, error) { + list, err := deps.Service.ListCoupons(ctx, scopedApp(deps, p)) + if err != nil { + return ackResponse{}, mapErr(err) + } + for _, item := range list { + if item.ID.String() == in.ID && item.AppID == scopedApp(deps, p) { + if err := deps.Service.DeleteCoupon(ctx, item.ID); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: in.ID}, nil + } + } + return ackResponse{}, notFound("coupon") + } +} + +type catalogFeatureInput struct { + ID string `json:"id"` + Key string `json:"key"` + Name string `json:"name"` + Description string `json:"description"` + Type string `json:"type"` + DefaultLimit int64 `json:"defaultLimit"` + Period string `json:"period"` + SoftLimit bool `json:"softLimit"` +} +type catalogFeatureSummary struct { + catalogFeatureInput + Status string `json:"status"` +} +type catalogFeaturesResponse struct { + Features []catalogFeatureSummary `json:"features"` +} + +func projectCatalogFeature(item *feature.Feature) catalogFeatureSummary { + return catalogFeatureSummary{catalogFeatureInput: catalogFeatureInput{ID: item.ID.String(), Key: item.Key, Name: item.Name, Description: item.Description, Type: string(item.Type), DefaultLimit: item.DefaultLimit, Period: string(item.Period), SoftLimit: item.SoftLimit}, Status: string(item.Status)} +} +func scopedCatalogFeature(ctx context.Context, deps Deps, p contract.Principal, raw string) (*feature.Feature, error) { + id, err := ledgerid.ParseFeatureID(strings.TrimSpace(raw)) + if err != nil { + return nil, badReq("invalid feature id") + } + item, err := deps.Service.GetCatalogFeature(ctx, id) + if err != nil { + return nil, mapErr(err) + } + if item == nil || item.AppID != scopedApp(deps, p) { + return nil, notFound("feature") + } + return item, nil +} +func validateCatalogFeature(in catalogFeatureInput) error { + if strings.TrimSpace(in.Key) == "" || strings.TrimSpace(in.Name) == "" { + return badReq("key and name are required") + } + if in.Type != string(feature.FeatureBoolean) && in.Type != string(feature.FeatureMetered) && in.Type != string(feature.FeatureSeat) { + return badReq("invalid feature type") + } + if in.Period != "" && in.Period != string(feature.PeriodNone) && in.Period != string(feature.PeriodMonthly) && in.Period != string(feature.PeriodYearly) { + return badReq("invalid feature period") + } + return nil +} +func featuresListHandler(deps Deps) func(context.Context, struct{}, contract.Principal) (catalogFeaturesResponse, error) { + return func(ctx context.Context, _ struct{}, p contract.Principal) (catalogFeaturesResponse, error) { + list, err := deps.Service.ListCatalogFeatures(ctx, scopedApp(deps, p)) + if err != nil { + return catalogFeaturesResponse{}, mapErr(err) + } + out := catalogFeaturesResponse{Features: make([]catalogFeatureSummary, 0, len(list))} + for _, item := range list { + if item.AppID == scopedApp(deps, p) { + out.Features = append(out.Features, projectCatalogFeature(item)) + } + } + return out, nil + } +} +func featuresCreateHandler(deps Deps) func(context.Context, catalogFeatureInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in catalogFeatureInput, p contract.Principal) (ackResponse, error) { + if err := validateCatalogFeature(in); err != nil { + return ackResponse{}, err + } + item := &feature.Feature{Key: strings.TrimSpace(in.Key), Name: strings.TrimSpace(in.Name), Description: in.Description, Type: feature.FeatureType(in.Type), DefaultLimit: in.DefaultLimit, Period: feature.Period(in.Period), SoftLimit: in.SoftLimit, Status: feature.StatusActive, AppID: scopedApp(deps, p)} + if err := deps.Service.CreateCatalogFeature(ctx, item); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func featuresUpdateHandler(deps Deps) func(context.Context, catalogFeatureInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in catalogFeatureInput, p contract.Principal) (ackResponse, error) { + item, err := scopedCatalogFeature(ctx, deps, p, in.ID) + if err != nil { + return ackResponse{}, err + } + if err := validateCatalogFeature(in); err != nil { + return ackResponse{}, err + } + if item.Key != strings.TrimSpace(in.Key) { + return ackResponse{}, badReq("feature key cannot change") + } + item.Name = strings.TrimSpace(in.Name) + item.Description = in.Description + item.Type = feature.FeatureType(in.Type) + item.DefaultLimit = in.DefaultLimit + item.Period = feature.Period(in.Period) + item.SoftLimit = in.SoftLimit + if err := deps.Service.UpdateCatalogFeature(ctx, item); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} +func featuresArchiveHandler(deps Deps) func(context.Context, catalogFeatureInput, contract.Principal) (ackResponse, error) { + return func(ctx context.Context, in catalogFeatureInput, p contract.Principal) (ackResponse, error) { + item, err := scopedCatalogFeature(ctx, deps, p, in.ID) + if err != nil { + return ackResponse{}, err + } + if err := deps.Service.ArchiveCatalogFeature(ctx, item.ID); err != nil { + return ackResponse{}, mapErr(err) + } + return ackResponse{OK: true, ID: item.ID.String()}, nil + } +} diff --git a/plugins/subscription/contract/handlers_test.go b/plugins/subscription/contract/handlers_test.go index ad6b90ea..c0c4fc00 100644 --- a/plugins/subscription/contract/handlers_test.go +++ b/plugins/subscription/contract/handlers_test.go @@ -3,7 +3,9 @@ package contract import ( "bytes" "context" + "encoding/json" "errors" + "strings" "testing" "github.com/xraph/forge/extensions/dashboard/contract" @@ -18,11 +20,8 @@ func TestManifest_Loads(t *testing.T) { if m.Contributor.Name != "subscription" { t.Errorf("contributor name = %q, want subscription", m.Contributor.Name) } - if got := len(m.Intents); got != 5 { - t.Errorf("intents = %d, want 5 (plans list/detail/archive/activate + subscriptions.list)", got) - } - if got := len(m.Graph); got != 2 { - t.Errorf("graph routes = %d, want 2 (/plans + /plans/:id)", got) + if got := len(m.Intents); got != 26 { + t.Errorf("intents = %d, want 26", got) } } @@ -44,3 +43,30 @@ func TestPlansListHandler_UnavailableWhenServiceNil(t *testing.T) { t.Errorf("expected CodeUnavailable, got %v", err) } } + +func TestPricingTiersRequirePlanFeature(t *testing.T) { + price, err := pricing(0, "usd", "monthly", []PriceTierSummary{{FeatureKey: "requests", Type: "graduated", UpTo: 1000, UnitAmount: 2}}) + if err != nil { + t.Fatalf("pricing: %v", err) + } + if err := validateTierFeatures(price, nil); err == nil { + t.Fatal("expected missing feature to be rejected") + } + features, err := planFeatures([]PlanFeature{{Key: "requests", Name: "Requests", Type: "metered", Limit: 1000, Period: "monthly"}}) + if err != nil { + t.Fatalf("features: %v", err) + } + if err := validateTierFeatures(price, features); err != nil { + t.Fatalf("valid tier rejected: %v", err) + } +} + +func TestInvoiceDetailWireIsFlat(t *testing.T) { + encoded, err := json.Marshal(invoiceDetail{invoiceSummary: invoiceSummary{ID: "inv_1", Status: "pending", Total: 1200}}) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(encoded), `"id":"inv_1"`) || !strings.Contains(string(encoded), `"total":1200`) { + t.Fatalf("invoice detail wire shape: %s", encoded) + } +} diff --git a/plugins/subscription/contract/manifest.yaml b/plugins/subscription/contract/manifest.yaml index aca50444..4b080759 100644 --- a/plugins/subscription/contract/manifest.yaml +++ b/plugins/subscription/contract/manifest.yaml @@ -16,6 +16,27 @@ intents: - { name: plans.archive, kind: command, version: 1, capability: write, invalidates: [plans.list, plans.detail] } - { name: plans.activate, kind: command, version: 1, capability: write, invalidates: [plans.list, plans.detail] } - { name: subscriptions.list, kind: query, version: 1, capability: read } + - { name: plans.create, kind: command, version: 1, capability: write, invalidates: [plans.list] } + - { name: plans.update, kind: command, version: 1, capability: write, invalidates: [plans.list, plans.detail] } + - { name: subscriptions.all, kind: query, version: 1, capability: read } + - { name: subscriptions.detail, kind: query, version: 1, capability: read } + - { name: subscriptions.create, kind: command, version: 1, capability: write, invalidates: [subscriptions.all, subscriptions.list] } + - { name: subscriptions.change, kind: command, version: 1, capability: write, invalidates: [subscriptions.all, subscriptions.list] } + - { name: subscriptions.pause, kind: command, version: 1, capability: write, invalidates: [subscriptions.all, subscriptions.list] } + - { name: subscriptions.resume, kind: command, version: 1, capability: write, invalidates: [subscriptions.all, subscriptions.list] } + - { name: subscriptions.cancel, kind: command, version: 1, capability: write, invalidates: [subscriptions.all, subscriptions.list] } + - { name: invoices.list, kind: query, version: 1, capability: read } + - { name: invoices.detail, kind: query, version: 1, capability: read } + - { name: invoices.generate, kind: command, version: 1, capability: write, invalidates: [invoices.list] } + - { name: invoices.markPaid, kind: command, version: 1, capability: write, invalidates: [invoices.list] } + - { name: invoices.void, kind: command, version: 1, capability: write, invalidates: [invoices.list] } + - { name: coupons.list, kind: query, version: 1, capability: read } + - { name: coupons.create, kind: command, version: 1, capability: write, invalidates: [coupons.list] } + - { name: coupons.delete, kind: command, version: 1, capability: write, invalidates: [coupons.list] } + - { name: features.list, kind: query, version: 1, capability: read } + - { name: features.create, kind: command, version: 1, capability: write, invalidates: [features.list] } + - { name: features.update, kind: command, version: 1, capability: write, invalidates: [features.list] } + - { name: features.archive, kind: command, version: 1, capability: write, invalidates: [features.list] } queries: planList: diff --git a/plugins/vpndetect/contract/handlers_test.go b/plugins/vpndetect/contract/handlers_test.go index 8232f749..d174ed74 100644 --- a/plugins/vpndetect/contract/handlers_test.go +++ b/plugins/vpndetect/contract/handlers_test.go @@ -19,9 +19,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 0 { t.Errorf("intents = %d, want 0", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/plugins/waitlist/contract/handlers_test.go b/plugins/waitlist/contract/handlers_test.go index e2fb37fc..63f81b64 100644 --- a/plugins/waitlist/contract/handlers_test.go +++ b/plugins/waitlist/contract/handlers_test.go @@ -21,9 +21,6 @@ func TestManifest_Loads(t *testing.T) { if got := len(m.Intents); got != 6 { t.Errorf("intents = %d, want 6 (list/detail/approve/reject/delete/counts)", got) } - if got := len(m.Graph); got != 1 { - t.Errorf("graph routes = %d, want 1 (/waitlist)", got) - } } func TestManifest_Validates(t *testing.T) { diff --git a/ui/packages/core/src/auth.ts b/ui/packages/core/src/auth.ts index 6818fc09..6b609ef9 100644 --- a/ui/packages/core/src/auth.ts +++ b/ui/packages/core/src/auth.ts @@ -112,6 +112,25 @@ function isAuthRejection(err: unknown): boolean { return false; } +// The browser lock covers tabs. The queue also covers managers sharing a +// storage adapter in runtimes without Web Locks. +const refreshQueues = new Map>(); + +async function withRefreshLock(key: string, operation: () => Promise): Promise { + if (typeof navigator !== "undefined" && navigator.locks) { + await navigator.locks.request(key, operation); + return; + } + const previous = refreshQueues.get(key) ?? Promise.resolve(); + const next = previous.catch(() => {}).then(operation); + refreshQueues.set(key, next); + try { + await next; + } finally { + if (refreshQueues.get(key) === next) refreshQueues.delete(key); + } +} + /** * AuthManager is the core state machine that drives authentication. * @@ -128,6 +147,10 @@ export class AuthManager { private storage: TokenStorage; private state: AuthState = { status: "idle" }; private listeners = new Set<(state: AuthState) => void>(); + private initializePromise: Promise | null = null; + private refreshPromise: Promise | null = null; + private refreshLockKey: string; + private active = true; private refreshTimer: ReturnType | null = null; private onError?: (error: { error: string; code?: number; type?: string }) => void; @@ -138,6 +161,7 @@ export class AuthManager { constructor(config: AuthConfig) { this.client = new AuthClient(config); + this.refreshLockKey = `authsome:refresh:${config.baseURL}:${config.publishableKey ?? ""}`; this.storage = config.storage ?? createMemoryStorage(); this.onError = config.onError; this.publishableKey = config.publishableKey; @@ -176,7 +200,18 @@ export class AuthManager { * When a publishableKey is set, also fetches client config in parallel. * Call this once on app start. */ - async initialize(): Promise { + initialize(): Promise { + this.active = true; + if (!this.initializePromise) { + this.initializePromise = this.restoreSession().finally(() => { + this.initializePromise = null; + }); + } + return this.initializePromise; + } + + private async restoreSession(): Promise { + let session: Session | undefined; // Kick off config fetch in parallel (non-blocking). if (this.publishableKey && !this.clientConfig) { void this.fetchClientConfig(); @@ -189,7 +224,7 @@ export class AuthManager { return; } - const session: Session = JSON.parse(raw); + session = JSON.parse(raw) as Session; const expiresAt = new Date(session.expires_at).getTime(); if (Date.now() >= expiresAt) { @@ -204,7 +239,11 @@ export class AuthManager { this.setState({ status: "authenticated", user, session }); this.scheduleRefresh(session); } catch (err) { - await this.handleInitFailure(err); + if (session?.refresh_token && isAuthRejection(err)) { + await this.refreshSession(session.refresh_token); + } else { + await this.handleInitFailure(err); + } } } @@ -245,7 +284,10 @@ export class AuthManager { } this.setState({ status: "unknown", session }); - this.scheduleRefresh(session); + this.clearRefreshTimer(); + if (this.active) { + this.refreshTimer = setTimeout(() => { void this.initialize(); }, REFRESH_RETRY_BASE_MS); + } } /** Sign in with email & password. */ @@ -447,9 +489,10 @@ export class AuthManager { /** Refresh the current session manually. */ async refreshNow(): Promise { - const state = this.state; - if (state.status !== "authenticated") return; - await this.refreshSession(state.session.refresh_token); + const raw = await this.storage.getItem(SESSION_KEY); + if (!raw) return; + const session = JSON.parse(raw) as Session; + if (session.refresh_token) await this.refreshSession(session.refresh_token); } /** @@ -545,6 +588,7 @@ export class AuthManager { /** Tear down: clear timers and listeners. */ destroy(): void { + this.active = false; this.clearRefreshTimer(); this.listeners.clear(); this.configListeners.clear(); @@ -558,62 +602,55 @@ export class AuthManager { this.scheduleRefresh(session); } - /** - * Exchanges the refresh token for a new session. - * - * On failure the two cases are separated, as they are in initialize: a - * rejected token means signed out, anything else means no verdict and is - * worth retrying. - * - * `attempt` was previously accepted and never read, and the retry path - * called back in without it — so a permanently-invalid token retried every - * 30s for as long as the tab stayed open, and the user was never signed out. - * It now bounds the retries and backs off between them. - */ - private async refreshSession(refreshToken: string, attempt = 0): Promise { + // Keep rotation and persistence in the same lock. Once the server rotates a + // token, every later operation must read its replacement from storage. + private refreshSession(refreshToken: string, attempt = 0): Promise { + if (!this.refreshPromise) { + this.refreshPromise = withRefreshLock(this.refreshLockKey, () => + this.exchangeSession(refreshToken, attempt), + ).finally(() => { this.refreshPromise = null; }); + } + return this.refreshPromise; + } + + private async exchangeSession(refreshToken: string, attempt: number): Promise { + let current: Session | undefined; + let exchanged = false; try { - const newSession = await this.client.refresh(refreshToken); - const user = await this.client.getMe(newSession.session_token); - await this.persistSession(newSession); - this.setState({ status: "authenticated", user, session: newSession }); - this.scheduleRefresh(newSession); - } catch (err) { - // The server rejected the token: retrying cannot help. - if (isAuthRejection(err)) { + const raw = await this.storage.getItem(SESSION_KEY); + if (!raw) { this.setState({ status: "unauthenticated" }); return; } - - if (attempt >= MAX_REFRESH_ATTEMPTS) { - // Out of retries without ever reaching the server. The session is - // retained — the token may still be good — but nothing has validated - // it, so we must not keep claiming authenticated. - const stored = await this.storage.getItem(SESSION_KEY); - if (stored) { - try { - this.setState({ status: "unknown", session: JSON.parse(stored) }); - return; - } catch { - // Fall through to unauthenticated on unparseable storage. - } - } + current = JSON.parse(raw) as Session; + // A waiting tab or an old timer may hold the token another tab spent. + if (current.refresh_token === refreshToken || Date.now() >= Date.parse(current.expires_at)) { + current = await this.client.refresh(current.refresh_token); + exchanged = true; + await this.persistSession(current); + } + const user = await this.client.getMe(current.session_token); + this.setState({ status: "authenticated", user, session: current }); + this.scheduleRefresh(current); + } catch (err) { + this.clearRefreshTimer(); + if (!exchanged && isAuthRejection(err) && current?.refresh_token === refreshToken) { + await this.clearSession(); this.setState({ status: "unauthenticated" }); return; } - - // No verdict — back off and try again. - this.clearRefreshTimer(); - this.refreshTimer = setTimeout( - () => { - void this.refreshSession(refreshToken, attempt + 1); - }, - REFRESH_RETRY_BASE_MS * 2 ** attempt, - ); + if (current) this.setState({ status: "unknown", session: current }); + if (attempt >= MAX_REFRESH_ATTEMPTS || !this.active) return; + this.refreshTimer = setTimeout(() => { + // Re-read storage under the lock, including after a profile failure. + void this.refreshSession(refreshToken, attempt + 1); + }, REFRESH_RETRY_BASE_MS * 2 ** attempt); } } private scheduleRefresh(session: Session): void { this.clearRefreshTimer(); + if (!this.active) return; const expiresAt = new Date(session.expires_at).getTime(); const delay = expiresAt - Date.now() - REFRESH_BEFORE_MS; diff --git a/ui/packages/core/src/client-config.test.ts b/ui/packages/core/src/client-config.test.ts new file mode 100644 index 00000000..059a21a9 --- /dev/null +++ b/ui/packages/core/src/client-config.test.ts @@ -0,0 +1,27 @@ +import { describe, expect, it, vi } from "vitest"; + +import { AuthClient } from "./client"; + +// Guards the base URL contract: an AuthSome API mounted under a path prefix +// (a gateway routes it at /identity/authsome) must keep that prefix on the +// client-config request. `new URL("/v1/...", base)` resolves against the +// origin and silently drops the path, so every page asked the gateway root +// and got a 404 plus an unhandled rejection. +describe("fetchClientConfig", () => { + it("keeps the base URL's path prefix and appends the publishable key", async () => { + const fetchFn = vi.fn((_input: string | URL | Request, _init?: RequestInit) => Promise.resolve(Response.json({ methods: [] }))); + const client = new AuthClient({ + baseURL: "https://gw.test/identity/authsome/", + publishableKey: "pk_x", + fetch: fetchFn, + }); + + await client.fetchClientConfig("pk_x"); + + expect(fetchFn).toHaveBeenCalledTimes(1); + expect(fetchFn.mock.calls[0]![0]).toBe("https://gw.test/identity/authsome/v1/client-config?key=pk_x"); + expect(fetchFn.mock.calls[0]![1]?.headers).toMatchObject({ + "X-Publishable-Key": "pk_x", + }); + }); +}); diff --git a/ui/packages/core/src/client.ts b/ui/packages/core/src/client.ts index 44430552..963fd36f 100644 --- a/ui/packages/core/src/client.ts +++ b/ui/packages/core/src/client.ts @@ -219,27 +219,8 @@ export class AuthClient extends GeneratedClient { * * The config describes which auth methods are enabled so SDK * components can auto-configure without manual props. - * - * Reads `baseURL` and `fetchFn` directly off the generated parent - * class — earlier generator templates kept a `this.config` ref but - * the current embedded template stores fields individually, so the - * old `(this as any).config.baseURL` shape produced `undefined` and - * the URL constructor threw. */ async fetchClientConfig(publishableKey?: string): Promise { - const baseURL = (this as any).baseURL as string; - const fetchFn = ((this as any).fetchFn as typeof globalThis.fetch) ?? globalThis.fetch; - const url = new URL("/v1/client-config", baseURL); - if (publishableKey) { - url.searchParams.set("key", publishableKey); - } - const res = await fetchFn(url.toString(), { - method: "GET", - headers: { "Content-Type": "application/json" }, - }); - if (!res.ok) { - throw new AuthClientError("Failed to fetch client config", res.status); - } - return res.json(); + return super.getClientConfig("", publishableKey); } } diff --git a/ui/packages/core/src/refresh-recovery.test.ts b/ui/packages/core/src/refresh-recovery.test.ts new file mode 100644 index 00000000..a37dd5fa --- /dev/null +++ b/ui/packages/core/src/refresh-recovery.test.ts @@ -0,0 +1,131 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { AuthManager } from "./auth"; +import { AuthClientError } from "./client"; +import type { Session, User } from "./types"; + +const user = { id: "user_1", email: "dev@example.test", roles: [] } as unknown as User & { roles: string[] }; +const oldSession: Session = { + session_token: "old-access", + refresh_token: "old-refresh", + expires_at: "2000-01-01T00:00:00Z", +}; +const freshSession = (): Session => ({ + session_token: "new-access", + refresh_token: "new-refresh", + expires_at: new Date(Date.now() + 3_600_000).toISOString(), +}); + +function setup(session = oldSession) { + const values = new Map([["authsome:session", JSON.stringify(session)]]); + const storage = { + getItem: (key: string) => values.get(key) ?? null, + setItem: (key: string, value: string) => { values.set(key, value); }, + removeItem: (key: string) => { values.delete(key); }, + }; + const manager = () => new AuthManager({ baseURL: "https://auth.test", storage }); + return { values, manager }; +} + +beforeEach(() => vi.useFakeTimers()); +afterEach(() => { vi.useRealTimers(); vi.unstubAllGlobals(); }); + +describe("session recovery", () => { + it("persists rotated tokens before a profile request can fail", async () => { + const { values, manager } = setup(); + const auth = manager(); + const replacement = freshSession(); + vi.spyOn(auth.getClient(), "refresh").mockResolvedValue(replacement); + const profile = vi.spyOn(auth.getClient(), "getMe") + .mockRejectedValueOnce(new TypeError("server restarting")) + .mockResolvedValue(user); + + await auth.initialize(); + expect(JSON.parse(values.get("authsome:session")!)).toEqual(replacement); + expect(auth.getState().status).toBe("unknown"); + await vi.advanceTimersByTimeAsync(30_000); + expect(auth.getState()).toMatchObject({ status: "authenticated", session: replacement }); + expect(profile).toHaveBeenLastCalledWith("new-access"); + expect(auth.getClient().refresh).toHaveBeenCalledTimes(1); + auth.destroy(); + }); + + it("shares initialization when React mounts an effect twice", async () => { + const { manager } = setup(); + const auth = manager(); + let exchanges = 0; + vi.spyOn(auth.getClient(), "refresh").mockImplementation(async () => { + if (++exchanges > 1) throw new AuthClientError("spent refresh token", 401); + return freshSession(); + }); + vi.spyOn(auth.getClient(), "getMe").mockResolvedValue(user); + const first = auth.initialize(); + auth.destroy(); + await Promise.all([first, auth.initialize()]); + expect(auth.getState().status).toBe("authenticated"); + expect(exchanges).toBe(1); + auth.destroy(); + }); + + it("adopts tokens saved by another manager before exchanging a stale token", async () => { + const { manager } = setup(); + const first = manager(); + const second = manager(); + let exchanges = 0; + const refresh = async () => { + if (++exchanges > 1) throw new AuthClientError("spent refresh token", 401); + return freshSession(); + }; + for (const auth of [first, second]) { + vi.spyOn(auth.getClient(), "refresh").mockImplementation(refresh); + vi.spyOn(auth.getClient(), "getMe").mockResolvedValue(user); + } + await Promise.all([first.initialize(), second.initialize()]); + expect(first.getState().status).toBe("authenticated"); + expect(second.getState().status).toBe("authenticated"); + expect(exchanges).toBe(1); + first.destroy(); second.destroy(); + }); + + it("uses the browser lock to serialize refresh across tabs", async () => { + const { manager } = setup(); + const auth = manager(); + let locked = false; + const request = vi.fn(async (_name: string, callback: () => Promise) => { + locked = true; + try { await callback(); } finally { locked = false; } + }); + vi.stubGlobal("navigator", { locks: { request } }); + vi.spyOn(auth.getClient(), "refresh").mockImplementation(async () => { + expect(locked).toBe(true); + return freshSession(); + }); + vi.spyOn(auth.getClient(), "getMe").mockResolvedValue(user); + await auth.initialize(); + expect(request).toHaveBeenCalledTimes(1); + expect(auth.getState().status).toBe("authenticated"); + auth.destroy(); + }); + + it("renews an access token rejected before its advertised expiry", async () => { + const { manager } = setup({ ...oldSession, expires_at: freshSession().expires_at }); + const auth = manager(); + vi.spyOn(auth.getClient(), "getMe") + .mockRejectedValueOnce(new AuthClientError("expired access token", 401)) + .mockResolvedValue(user); + vi.spyOn(auth.getClient(), "refresh").mockResolvedValue(freshSession()); + await auth.initialize(); + expect(auth.getState().status).toBe("authenticated"); + expect(auth.getClient().refresh).toHaveBeenCalledTimes(1); + auth.destroy(); + }); +}); + +it("removes rejected credentials so recovery cannot loop back from sign-in", async () => { + const { values, manager } = setup(); + const auth = manager(); + vi.spyOn(auth.getClient(), "refresh").mockRejectedValue(new AuthClientError("revoked", 401)); + await auth.initialize(); + expect(auth.getState().status).toBe("unauthenticated"); + expect(values.has("authsome:session")).toBe(false); + auth.destroy(); +}); diff --git a/ui/packages/nextjs/src/client-config.test.ts b/ui/packages/nextjs/src/client-config.test.ts new file mode 100644 index 00000000..40334c4b --- /dev/null +++ b/ui/packages/nextjs/src/client-config.test.ts @@ -0,0 +1,23 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; + +import { getClientConfig } from "./server"; + +afterEach(() => vi.unstubAllGlobals()); + +// Same contract as the browser client: a path-prefixed base URL keeps its +// prefix. The server helper is what a Next app calls at render time to seed +// the provider, so a wrong URL here means a null config on every request. +describe("getClientConfig", () => { + it("keeps the base URL's path prefix and appends the publishable key", async () => { + const fetchFn = vi.fn((_input: string | URL | Request, _init?: RequestInit) => Promise.resolve(Response.json({ methods: [] }))); + vi.stubGlobal("fetch", fetchFn); + + const config = await getClientConfig({ baseURL: "https://gw.test/identity/authsome", publishableKey: "pk_x" }); + + expect(config).toEqual({ methods: [] }); + expect(fetchFn.mock.calls[0]![0]).toBe("https://gw.test/identity/authsome/v1/client-config?key=pk_x"); + expect(fetchFn.mock.calls[0]![1]?.headers).toMatchObject({ + "X-Publishable-Key": "pk_x", + }); + }); +}); diff --git a/ui/packages/nextjs/src/server.ts b/ui/packages/nextjs/src/server.ts index 495e2bcb..a2954b4c 100644 --- a/ui/packages/nextjs/src/server.ts +++ b/ui/packages/nextjs/src/server.ts @@ -95,19 +95,17 @@ export interface GetClientConfigOptions { export async function getClientConfig( opts: GetClientConfigOptions, ): Promise { - const url = new URL("/v1/client-config", opts.baseURL); - if (opts.publishableKey) { - url.searchParams.set("key", opts.publishableKey); - } - try { - const fetchOpts: RequestInit & Record = { - headers: { "Content-Type": "application/json" }, - next: { revalidate: 300 }, - }; - const res = await fetch(url.toString(), fetchOpts as RequestInit); - if (!res.ok) return null; - return (await res.json()) as ClientConfig; + const client = new AuthClient({ + baseURL: opts.baseURL, + publishableKey: opts.publishableKey, + fetch: (input, init) => + fetch(input, { + ...init, + next: { revalidate: 300 }, + } as RequestInit), + }); + return await client.getClientConfig("", opts.publishableKey); } catch { return null; }