From 33d694b89b568e66226f6329cef287c3e432ea4d Mon Sep 17 00:00:00 2001 From: Heatherm Huang Date: Thu, 23 Jul 2026 00:08:23 +0800 Subject: [PATCH] fix(simple-mode): enable images for auto Grok default --- ...he_invalidation_outbox_integration_test.go | 4 + .../repository/simple_mode_default_groups.go | 27 +++++- ...le_mode_default_groups_integration_test.go | 94 +++++++++++++++++++ .../186_group_auth_cache_image_generation.sql | 31 ++++++ 4 files changed, 155 insertions(+), 1 deletion(-) create mode 100644 backend/migrations/186_group_auth_cache_image_generation.sql diff --git a/backend/internal/repository/auth_cache_invalidation_outbox_integration_test.go b/backend/internal/repository/auth_cache_invalidation_outbox_integration_test.go index 5e8cc286f2..eeb46e44e2 100644 --- a/backend/internal/repository/auth_cache_invalidation_outbox_integration_test.go +++ b/backend/internal/repository/auth_cache_invalidation_outbox_integration_test.go @@ -92,6 +92,10 @@ func TestAuthCacheInvalidationTriggers_CoverSecurityMutationsOnly(t *testing.T) _, err = integrationDB.ExecContext(ctx, "UPDATE groups SET name = name || '-cosmetic' WHERE id = $1", group.ID) require.NoError(t, err) require.Zero(t, count(), "cosmetic group update must not enqueue") + _, err = integrationDB.ExecContext(ctx, "UPDATE groups SET allow_image_generation = NOT allow_image_generation WHERE id = $1", group.ID) + require.NoError(t, err) + require.Equal(t, 1, count(), "image-generation permission changes must enqueue bound keys") + clear() _, err = integrationDB.ExecContext(ctx, "UPDATE groups SET status = 'disabled' WHERE id = $1", group.ID) require.NoError(t, err) require.Equal(t, 1, count(), "group disable must enqueue bound keys") diff --git a/backend/internal/repository/simple_mode_default_groups.go b/backend/internal/repository/simple_mode_default_groups.go index e3786451a4..2b9daa4d64 100644 --- a/backend/internal/repository/simple_mode_default_groups.go +++ b/backend/internal/repository/simple_mode_default_groups.go @@ -9,11 +9,17 @@ import ( "github.com/Wei-Shaw/sub2api/internal/service" ) +const simpleModeDefaultGroupDescription = "Auto-created default group" + func ensureSimpleModeDefaultGroups(ctx context.Context, client *dbent.Client) error { if client == nil { return fmt.Errorf("nil ent client") } + if err := backfillSimpleModeGrokDefaultImageGeneration(ctx, client); err != nil { + return err + } + requiredByPlatform := map[string]int{ service.PlatformAnthropic: 1, service.PlatformOpenAI: 1, @@ -65,12 +71,13 @@ func createGroupIfNotExists(ctx context.Context, client *dbent.Client, name, pla _, err = client.Group.Create(). SetName(name). - SetDescription("Auto-created default group"). + SetDescription(simpleModeDefaultGroupDescription). SetPlatform(platform). SetStatus(service.StatusActive). SetSubscriptionType(service.SubscriptionTypeStandard). SetRateMultiplier(1.0). SetIsExclusive(false). + SetAllowImageGeneration(platform == service.PlatformGrok). Save(ctx) if err != nil { if dbent.IsConstraintError(err) { @@ -81,3 +88,21 @@ func createGroupIfNotExists(ctx context.Context, client *dbent.Client, name, pla } return nil } + +func backfillSimpleModeGrokDefaultImageGeneration(ctx context.Context, client *dbent.Client) error { + _, err := client.Group.Update(). + Where( + group.NameEQ(service.PlatformGrok+"-default"), + group.PlatformEQ(service.PlatformGrok), + group.DescriptionEQ(simpleModeDefaultGroupDescription), + group.StatusEQ(service.StatusActive), + group.AllowImageGenerationEQ(false), + group.DeletedAtIsNil(), + ). + SetAllowImageGeneration(true). + Save(ctx) + if err != nil { + return fmt.Errorf("backfill auto-created grok default image generation: %w", err) + } + return nil +} diff --git a/backend/internal/repository/simple_mode_default_groups_integration_test.go b/backend/internal/repository/simple_mode_default_groups_integration_test.go index 3327257b40..d6a96bbfa2 100644 --- a/backend/internal/repository/simple_mode_default_groups_integration_test.go +++ b/backend/internal/repository/simple_mode_default_groups_integration_test.go @@ -33,6 +33,100 @@ func TestEnsureSimpleModeDefaultGroups_CreatesMissingDefaults(t *testing.T) { assertGroupExists(service.PlatformGemini + "-default") assertGroupExists(service.PlatformAntigravity + "-default-1") assertGroupExists(service.PlatformAntigravity + "-default-2") + + grokDefault, err := client.Group.Query(). + Where(group.NameEQ(service.PlatformGrok+"-default"), group.DeletedAtIsNil()). + Only(seedCtx) + require.NoError(t, err) + require.True(t, grokDefault.AllowImageGeneration) +} + +func TestEnsureSimpleModeDefaultGroups_BackfillsOnlyAutoCreatedGrokDefault(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + client := tx.Client() + + seedCtx, cancel := context.WithTimeout(ctx, 10*time.Second) + defer cancel() + + autoDefault, err := client.Group.Create(). + SetName(service.PlatformGrok + "-default"). + SetDescription("Auto-created default group"). + SetPlatform(service.PlatformGrok). + SetStatus(service.StatusActive). + SetSubscriptionType(service.SubscriptionTypeStandard). + SetRateMultiplier(1.0). + SetIsExclusive(false). + SetAllowImageGeneration(false). + Save(seedCtx) + require.NoError(t, err) + + operatorGroup, err := client.Group.Create(). + SetName("operator-grok-images-disabled-" + time.Now().Format(time.RFC3339Nano)). + SetDescription("Operator-managed group"). + SetPlatform(service.PlatformGrok). + SetStatus(service.StatusActive). + SetSubscriptionType(service.SubscriptionTypeStandard). + SetRateMultiplier(1.0). + SetIsExclusive(false). + SetAllowImageGeneration(false). + Save(seedCtx) + require.NoError(t, err) + + require.NoError(t, ensureSimpleModeDefaultGroups(seedCtx, client)) + + autoDefault, err = client.Group.Get(seedCtx, autoDefault.ID) + require.NoError(t, err) + require.True(t, autoDefault.AllowImageGeneration) + + operatorGroup, err = client.Group.Get(seedCtx, operatorGroup.ID) + require.NoError(t, err) + require.False(t, operatorGroup.AllowImageGeneration, "operator-managed false must be preserved") +} + +func TestEnsureSimpleModeDefaultGroups_PreservesExplicitFalse(t *testing.T) { + tests := []struct { + name string + description string + status string + }{ + { + name: "operator managed default", + description: "Operator-managed group", + status: service.StatusActive, + }, + { + name: "disabled auto-created default", + description: simpleModeDefaultGroupDescription, + status: service.StatusDisabled, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + client := testEntTx(t).Client() + grokDefault, err := client.Group.Create(). + SetName(service.PlatformGrok + "-default"). + SetDescription(tt.description). + SetPlatform(service.PlatformGrok). + SetStatus(tt.status). + SetSubscriptionType(service.SubscriptionTypeStandard). + SetRateMultiplier(1.0). + SetIsExclusive(false). + SetAllowImageGeneration(false). + Save(ctx) + require.NoError(t, err) + + require.NoError(t, ensureSimpleModeDefaultGroups(ctx, client)) + + grokDefault, err = client.Group.Get(ctx, grokDefault.ID) + require.NoError(t, err) + require.False(t, grokDefault.AllowImageGeneration) + }) + } } func TestEnsureSimpleModeDefaultGroups_IgnoresSoftDeletedGroups(t *testing.T) { diff --git a/backend/migrations/186_group_auth_cache_image_generation.sql b/backend/migrations/186_group_auth_cache_image_generation.sql new file mode 100644 index 0000000000..2c27a957a3 --- /dev/null +++ b/backend/migrations/186_group_auth_cache_image_generation.sql @@ -0,0 +1,31 @@ +-- Group image-generation permission is part of the API-key auth snapshot. +-- Extend the existing durable invalidation trigger without changing migration 184. + +CREATE OR REPLACE FUNCTION enqueue_group_auth_cache_invalidation() +RETURNS TRIGGER +LANGUAGE plpgsql +AS $$ +DECLARE + target_group_id BIGINT; +BEGIN + target_group_id := OLD.id; + IF TG_OP = 'UPDATE' + AND OLD.status IS NOT DISTINCT FROM NEW.status + AND OLD.is_exclusive IS NOT DISTINCT FROM NEW.is_exclusive + AND OLD.allow_image_generation IS NOT DISTINCT FROM NEW.allow_image_generation + AND OLD.deleted_at IS NOT DISTINCT FROM NEW.deleted_at THEN + RETURN NEW; + END IF; + + INSERT INTO auth_cache_invalidation_outbox (cache_key) + SELECT encode(sha256(convert_to(k.key, 'UTF8')), 'hex') + FROM api_keys AS k + WHERE k.group_id = target_group_id + AND k.deleted_at IS NULL + AND k.key <> ''; + IF TG_OP = 'DELETE' THEN + RETURN OLD; + END IF; + RETURN NEW; +END; +$$;