From 9c2e7a748697d1b2a5256f7728d119eb5ba3c923 Mon Sep 17 00:00:00 2001 From: Ethan Dickson Date: Sun, 9 Aug 2026 06:05:24 +0000 Subject: [PATCH 1/3] feat: add organization chat model migration and telemetry --- ...6_chat_model_config_org_explosion.down.sql | 113 ++++ ...566_chat_model_config_org_explosion.up.sql | 178 ++++++ coderd/database/migrations/migrate.go | 8 + coderd/database/migrations/migrate_test.go | 535 ++++++++++++++++++ ...566_chat_model_config_org_explosion.up.sql | 33 ++ coderd/database/queries.sql.go | 16 +- coderd/database/queries/chats.sql | 2 +- coderd/telemetry/telemetry.go | 26 +- coderd/telemetry/telemetry_test.go | 2 + 9 files changed, 893 insertions(+), 20 deletions(-) create mode 100644 coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql create mode 100644 coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql create mode 100644 coderd/database/migrations/testdata/fixtures/000566_chat_model_config_org_explosion.up.sql diff --git a/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql b/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql new file mode 100644 index 00000000000..27e3f453c4e --- /dev/null +++ b/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql @@ -0,0 +1,113 @@ +-- DOWN for the chat model config org explosion. UNSUPPORTED and best-effort +-- per operator ruling: there is no persisted provenance, so this down +-- cannot distinguish a copy from an organically created non-default-org row +-- that happens to share (ai_provider_id, model) with a default-org row. It +-- must run green and must never lose chats; fidelity loss on pathological +-- duplicates is accepted. +-- +-- Copy identification: a non-default-org row is treated as a copy iff a +-- default-org row exists with the same (ai_provider_id, model). Retarget +-- resolves that default-org row deterministically with DISTINCT ON ordered +-- by (created_at ASC, id ASC) so duplicates pick one stable row. + +-- Restore chats.last_model_config_id from copies back to the default-org +-- original matched by (ai_provider_id, model). +UPDATE chats c +SET last_model_config_id = orig.id +FROM chat_model_configs cp +JOIN LATERAL ( + SELECT d.id + FROM chat_model_configs d + JOIN organizations def ON def.id = d.organization_id AND def.is_default + WHERE d.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id + AND d.model = cp.model + ORDER BY d.created_at ASC, d.id ASC + LIMIT 1 +) orig ON true +WHERE c.last_model_config_id = cp.id + AND NOT EXISTS (SELECT 1 FROM organizations odef + WHERE odef.id = cp.organization_id AND odef.is_default); + +-- Restore chat_messages.model_config_id from copies back to originals. +UPDATE chat_messages mm +SET model_config_id = orig.id +FROM chat_model_configs cp +JOIN LATERAL ( + SELECT d.id + FROM chat_model_configs d + JOIN organizations def ON def.id = d.organization_id AND def.is_default + WHERE d.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id + AND d.model = cp.model + ORDER BY d.created_at ASC, d.id ASC + LIMIT 1 +) orig ON true +WHERE mm.model_config_id = cp.id + AND NOT EXISTS (SELECT 1 FROM organizations odef + WHERE odef.id = cp.organization_id AND odef.is_default); + +-- Restore chat_queued_messages.model_config_id from copies back to originals. +UPDATE chat_queued_messages q +SET model_config_id = orig.id +FROM chat_model_configs cp +JOIN LATERAL ( + SELECT d.id + FROM chat_model_configs d + JOIN organizations def ON def.id = d.organization_id AND def.is_default + WHERE d.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id + AND d.model = cp.model + ORDER BY d.created_at ASC, d.id ASC + LIMIT 1 +) orig ON true +WHERE q.model_config_id = cp.id + AND NOT EXISTS (SELECT 1 FROM organizations odef + WHERE odef.id = cp.organization_id AND odef.is_default); + +-- Restore chat_debug_runs.model_config_id from copies back to originals. +UPDATE chat_debug_runs d +SET model_config_id = orig.id +FROM chat_model_configs cp +JOIN LATERAL ( + SELECT dc.id + FROM chat_model_configs dc + JOIN organizations def ON def.id = dc.organization_id AND def.is_default + WHERE dc.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id + AND dc.model = cp.model + ORDER BY dc.created_at ASC, dc.id ASC + LIMIT 1 +) orig ON true +WHERE d.model_config_id = cp.id + AND NOT EXISTS (SELECT 1 FROM organizations odef + WHERE odef.id = cp.organization_id AND odef.is_default); + +-- Delete copied chat_model_configs (non-default-org rows whose +-- (ai_provider_id, model) matches a default-org row). References were +-- retargeted above, so the deletes cannot violate the chats/chat_messages +-- FKs. +DELETE FROM chat_model_configs cp +WHERE EXISTS ( + SELECT 1 FROM chat_model_configs orig + JOIN organizations def ON def.id = orig.organization_id AND def.is_default + WHERE orig.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id + AND orig.model = cp.model +) +AND NOT EXISTS ( + SELECT 1 FROM organizations odef + WHERE odef.id = cp.organization_id AND odef.is_default +); + +-- Best-effort threshold-key cleanup: delete compaction-threshold keys whose +-- embedded config id no longer exists anywhere after the copy deletes. +-- Original keys survive because default-org originals always survive. +-- Keys with a malformed or empty suffix are guarded BEFORE the uuid cast +-- (they cannot name an existing config, so they are pruned like any other +-- dangling key) instead of aborting the down. +DELETE FROM user_configs uc +WHERE uc.key LIKE 'chat_compaction_threshold_pct:%' + AND NOT EXISTS ( + SELECT 1 FROM chat_model_configs cmc + WHERE cmc.id = ( + SELECT substring(uc.key FROM 'chat_compaction_threshold_pct:(.*)')::uuid + WHERE substring(uc.key FROM 'chat_compaction_threshold_pct:(.*)') + ~ '^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' + ) + ); diff --git a/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql b/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql new file mode 100644 index 00000000000..6476d72047a --- /dev/null +++ b/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql @@ -0,0 +1,178 @@ +-- Explode default-org chat model configs to every live non-default +-- organization (CODAGT-709, stage 3 of 3: org-scoping cutover). After this +-- migration every live org owns a full set of model configs and all +-- references inside live non-default orgs point at same-org rows. +-- +-- Mapping design (operator ruling): NO provenance column, NO persisted +-- mapping of any kind. A transaction-scoped TEMPORARY lookup table +-- (orig_id, org_id, copy_id) ON COMMIT DROP maps each default-org original +-- to its per-org copy; the copy insert, all four reference remaps, and the +-- compaction-threshold fan-out join it. The table vanishes when the +-- migration framework commits this migration's transaction, so nothing +-- mapping-related persists. Copy ids come from gen_random_uuid(); no +-- hash-derived ids (md5() errors on FIPS-mode PostgreSQL builds). + +CREATE TEMPORARY TABLE model_config_copy_map ( + orig_id uuid NOT NULL, + org_id uuid NOT NULL, + copy_id uuid NOT NULL, + PRIMARY KEY (orig_id, org_id) +) ON COMMIT DROP; + +-- (a) Stage LIVE default-org chat_model_configs x every live non-default org +-- in the temp map with a fresh copy id, then insert the copies. Staging the +-- id in the map first lets the remap statements below resolve copies without +-- recomputing anything. +INSERT INTO model_config_copy_map (orig_id, org_id, copy_id) +SELECT cmc.id, o.id, gen_random_uuid() +FROM chat_model_configs cmc +JOIN organizations def ON def.id = cmc.organization_id AND def.is_default +CROSS JOIN organizations o +WHERE NOT o.is_default AND NOT o.deleted + AND NOT cmc.deleted; + +-- Copies inherit every behavioral field from the original, including +-- created_at/updated_at and created_by/updated_by: a copy is the same +-- logical config re-homed, and the audit-facing identity of who configured +-- it survives the explosion. group_acl is re-keyed to the copy's org (the +-- Everyone group of an organization always has the organization's own ID, +-- see 000058) carrying the original's entry verbatim, so members of the +-- target org keep read access through the everyone entry. +INSERT INTO chat_model_configs + (id, model, display_name, created_by, updated_by, enabled, is_default, + deleted, deleted_at, created_at, updated_at, context_limit, + compression_threshold, options, ai_provider_id, organization_id, + group_acl, user_acl) +SELECT + m.copy_id, + cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, + cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, + cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, + cmc.ai_provider_id, m.org_id, + jsonb_build_object( + m.org_id::text, + COALESCE(cmc.group_acl -> cmc.organization_id::text, + '{"permissions": ["read"]}'::jsonb) + ), + '{}'::jsonb +FROM model_config_copy_map m +JOIN chat_model_configs cmc ON cmc.id = m.orig_id +WHERE NOT cmc.deleted; + +-- (a2) Stage + copy SOFT-DELETED default-org chat_model_configs ONLY to live +-- non-default orgs that actually reference them. A reference is any of: +-- chats.last_model_config_id, chat_messages.model_config_id (via chat), +-- chat_queued_messages.model_config_id (via chat), or +-- chat_debug_runs.model_config_id (via chat) pointing at the deleted config. +-- Copies keep deleted/deleted_at so every historical reference has an +-- FK-valid, attribution-preserving target without resurrecting the config. +INSERT INTO model_config_copy_map (orig_id, org_id, copy_id) +SELECT DISTINCT ON (cmc.id, o.id) cmc.id, o.id, gen_random_uuid() +FROM chat_model_configs cmc +JOIN organizations def ON def.id = cmc.organization_id AND def.is_default +JOIN organizations o ON NOT o.is_default AND NOT o.deleted +WHERE cmc.deleted + AND ( + EXISTS (SELECT 1 FROM chats c + WHERE c.last_model_config_id = cmc.id AND c.organization_id = o.id) + OR + EXISTS (SELECT 1 FROM chat_messages mm + JOIN chats c ON c.id = mm.chat_id + WHERE mm.model_config_id = cmc.id AND c.organization_id = o.id) + OR + EXISTS (SELECT 1 FROM chat_queued_messages q + JOIN chats c ON c.id = q.chat_id + WHERE q.model_config_id = cmc.id AND c.organization_id = o.id) + OR + EXISTS (SELECT 1 FROM chat_debug_runs d + JOIN chats c ON c.id = d.chat_id + WHERE d.model_config_id = cmc.id AND c.organization_id = o.id) + ); + +INSERT INTO chat_model_configs + (id, model, display_name, created_by, updated_by, enabled, is_default, + deleted, deleted_at, created_at, updated_at, context_limit, + compression_threshold, options, ai_provider_id, organization_id, + group_acl, user_acl) +SELECT + m.copy_id, + cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, + cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, + cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, + cmc.ai_provider_id, m.org_id, + jsonb_build_object( + m.org_id::text, + COALESCE(cmc.group_acl -> cmc.organization_id::text, + '{"permissions": ["read"]}'::jsonb) + ), + '{}'::jsonb +FROM model_config_copy_map m +JOIN chat_model_configs cmc ON cmc.id = m.orig_id +WHERE cmc.deleted; + +-- (b) Remap chats.last_model_config_id in live non-default orgs to the +-- same-org copy via the temp map. Soft-deleted orgs have no map rows, so +-- their chats keep original references. +UPDATE chats c +SET last_model_config_id = m.copy_id +FROM model_config_copy_map m +WHERE c.last_model_config_id = m.orig_id + AND m.org_id = c.organization_id; + +-- (b2) Remap chat_messages.model_config_id via the owning chat's org. +UPDATE chat_messages mm +SET model_config_id = m.copy_id +FROM chats c, model_config_copy_map m +WHERE c.id = mm.chat_id + AND mm.model_config_id = m.orig_id + AND m.org_id = c.organization_id; + +-- (b3) Remap chat_queued_messages.model_config_id via the owning chat's org. +-- The column has no FK, so dangling ids would not fail. The remap keeps a +-- queued message's promoted model inside its chat's org. +UPDATE chat_queued_messages q +SET model_config_id = m.copy_id +FROM chats c, model_config_copy_map m +WHERE c.id = q.chat_id + AND q.model_config_id = m.orig_id + AND m.org_id = c.organization_id; + +-- (b4) Remap chat_debug_runs.model_config_id via the owning chat's org. +-- The column is FK-less and stores attribution only. +UPDATE chat_debug_runs d +SET model_config_id = m.copy_id +FROM chats c, model_config_copy_map m +WHERE c.id = d.chat_id + AND d.model_config_id = m.orig_id + AND m.org_id = c.organization_id; + +-- (c) Fan out user_configs compaction-threshold keys. A key +-- 'chat_compaction_threshold_pct:' earns one row per copy of that +-- original in the temp map, same user, same value, key rewritten to the +-- copy id. The fan-out is copy-precise by construction (it can only +-- produce keys for copies that exist): live originals reach every live +-- org, soft-deleted originals reach only the orgs that received a +-- referenced copy, and an original with zero map rows (deleted and +-- unreferenced) produces nothing. Original keys stay: they reference +-- default-org originals, still valid. The fan-out is deliberately NOT +-- membership-filtered: chats pinned to deleted models are the norm, and a +-- threshold must keep resolving for any chat that lands on a copy. The PK +-- (user_id, key) cannot collide because copy ids are fresh and no existing +-- key embeds a copy id; ON CONFLICT DO NOTHING is belt-and-braces only. +INSERT INTO user_configs (user_id, key, value) +SELECT uc.user_id, 'chat_compaction_threshold_pct:' || m.copy_id::text, uc.value +FROM user_configs uc +JOIN model_config_copy_map m + ON uc.key = 'chat_compaction_threshold_pct:' || m.orig_id::text +ON CONFLICT (user_id, key) DO NOTHING; + +-- (d) Seed the everyone-in-org read entry on any existing row whose +-- group_acl lacks its own org's key. This covers rows written by older +-- binaries during a rolling upgrade. The entry's permissions are preserved +-- when an entry already exists for another org's key shape. +UPDATE chat_model_configs +SET group_acl = jsonb_build_object( + organization_id::text, + jsonb_build_object('permissions', jsonb_build_array('read'::text)) +) || group_acl +WHERE NOT (group_acl ? organization_id::text); diff --git a/coderd/database/migrations/migrate.go b/coderd/database/migrations/migrate.go index 50a931c902f..ec03a606f66 100644 --- a/coderd/database/migrations/migrate.go +++ b/coderd/database/migrations/migrate.go @@ -23,6 +23,14 @@ import ( //go:embed *.sql var migrations embed.FS +// MigrationFS exposes the embedded migration files, for tests that need to +// execute a single migration's SQL outside the migrate driver (the driver's +// transaction commits only when a stepper exhausts, which mid-test +// down-then-up cycles cannot wait for). +func MigrationFS() fs.FS { + return migrations +} + var ( migrationsHash string migrationsHashOnce sync.Once diff --git a/coderd/database/migrations/migrate_test.go b/coderd/database/migrations/migrate_test.go index 1f6a1b0b56f..a3ec5435267 100644 --- a/coderd/database/migrations/migrate_test.go +++ b/coderd/database/migrations/migrate_test.go @@ -5,6 +5,7 @@ import ( "database/sql" "encoding/json" "fmt" + "io/fs" "os" "path/filepath" "slices" @@ -2953,3 +2954,537 @@ func mustJSON(t *testing.T, v any) []byte { require.NoError(t, err) return raw } + +func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { + t.Parallel() + + const migrationVersion = 566 + + sqlDB := testSQLDB(t) + + // stepperUpToLatest runs the stepper to completion: a Stepper cannot + // stop early, it closes only when the steps are exhausted, and each + // call commits the driver's transaction. The assertion only requires + // that target was APPLIED, not that it is the latest: a stacked PR may + // add a later migration, so encoding "target is latest" would redden + // any such child on its merge ref. + stepperUpToLatest := func(target uint) { + t.Helper() + next, err := migrations.Stepper(sqlDB) + require.NoError(t, err) + last := uint(0) + for { + version, more, err := next() + require.NoError(t, err) + if !more { + break + } + last = version + } + require.GreaterOrEqual(t, last, target, "migration %d must be applied", target) + } + + // Migrate fully up first. The Stepper cannot stop at an intermediate + // version (it runs to completion and each run is one transaction), so + // fixtures are seeded post-schema and the explosion is exercised as + // DOWN (restore pre-explosion state) then UP (re-explode from it). + stepperUpToLatest(migrationVersion) + + ctx := testutil.Context(t, testutil.WaitSuperLong) + + now := time.Now().UTC().Truncate(time.Microsecond) + providerID := uuid.New() + user1ID := uuid.New() + user2ID := uuid.New() + orgBID := uuid.New() // live org with chats + orgCID := uuid.New() // live zero-member org: receives the full live set + orgDID := uuid.New() // soft-deleted org: receives nothing, chats untouched + c1ID := uuid.New() // live default config + c2ID := uuid.New() // live plain config + c3ID := uuid.New() // soft-deleted, referenced in orgB only + c4ID := uuid.New() // soft-deleted, unreferenced: never copied + c5ID := uuid.New() // live plain config + // aclLateConfigID simulates a config written by an older binary during a + // rolling upgrade, without the everyone ACL entry. + aclLateConfigID := uuid.New() + chatBID := uuid.New() // chat in orgB pinned to live c1 + chatB3ID := uuid.New() // chat in orgB pinned to deleted c3 + chatDID := uuid.New() // chat in soft-deleted orgD pinned to live c1 + + execFixture := func(query string, args ...any) { + t.Helper() + _, err := sqlDB.ExecContext(ctx, query, args...) + require.NoError(t, err) + } + + for i, id := range []uuid.UUID{user1ID, user2ID} { + execFixture( + `INSERT INTO users (id, username, email, hashed_password, created_at, updated_at, status, rbac_roles, login_type) + VALUES ($1, $2, $3, $4, $5, $6, 'active', '{}', 'password')`, + id, fmt.Sprintf("m3user%d", i+1), fmt.Sprintf("m3user%d@coder.com", i+1), []byte{}, now, now, + ) + } + + execFixture( + `INSERT INTO ai_providers (id, type, name, enabled, base_url, created_at, updated_at) + VALUES ($1, $2, $3, $4, $5, $6, $7)`, + providerID, "openai", "openai-566", true, "https://api.openai.com/v1", now, now, + ) + + // Three non-default orgs: live B, live zero-member C, soft-deleted D. + for _, o := range []struct { + id uuid.UUID + name string + deleted bool + }{ + {orgBID, "org-b-566", false}, + {orgCID, "org-c-566", false}, + {orgDID, "org-d-566", true}, + } { + execFixture( + `INSERT INTO organizations (id, name, description, display_name, default_org_member_roles, created_at, updated_at, deleted) + VALUES ($1, $2, '', '', '{}', $3, $3, $4)`, + o.id, o.name, now, o.deleted, + ) + } + + var defaultOrgID uuid.UUID + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT id FROM organizations WHERE is_default = true").Scan(&defaultOrgID)) + + // Model configs in the default org (pre-566 state: every config lives + // there, backfilled by 000565 with the everyone ACL entry). + insertConfig := func(id uuid.UUID, model string, isDefault, deleted bool, acl string) { + t.Helper() + execFixture( + `INSERT INTO chat_model_configs (id, model, display_name, enabled, is_default, deleted, deleted_at, + context_limit, compression_threshold, ai_provider_id, organization_id, group_acl, created_at, updated_at) + VALUES ($1, $2, $3, true, $4, $5, (CASE WHEN $5 THEN $6::timestamptz ELSE NULL END), + 200000, 70, $7, $8, $9::jsonb, $6, $6)`, + id, model, model+" display", isDefault, deleted, now, providerID, defaultOrgID, acl, + ) + } + everyoneACL := `{"` + defaultOrgID.String() + `": {"permissions": ["read"]}}` + insertConfig(c1ID, "gpt-5.2", true, false, everyoneACL) + insertConfig(c2ID, "gpt-5.2-mini", false, false, everyoneACL) + insertConfig(c3ID, "gpt-4-legacy", false, true, everyoneACL) + insertConfig(c4ID, "gpt-4-ancient", false, true, everyoneACL) + insertConfig(c5ID, "gpt-5.2-nano", false, false, everyoneACL) + insertConfig(aclLateConfigID, "gpt-5.2-late", false, false, `{}`) + + // Chats: orgB pinned to live c1, orgB second chat pinned to deleted c3, + // soft-deleted orgD pinned to live c1. + for _, ch := range []struct { + id uuid.UUID + orgID uuid.UUID + ownerID uuid.UUID + cfgID uuid.UUID + }{ + {chatBID, orgBID, user1ID, c1ID}, + {chatB3ID, orgBID, user2ID, c3ID}, + {chatDID, orgDID, user1ID, c1ID}, + } { + execFixture( + `INSERT INTO chats (id, owner_id, organization_id, last_model_config_id, created_at, updated_at) + VALUES ($1, $2, $3, $4, $5, $5)`, + ch.id, ch.ownerID, ch.orgID, ch.cfgID, now, + ) + } + + // Messages in each chat referencing the chat's pinned config. + for _, m := range []struct { + chatID uuid.UUID + cfgID uuid.UUID + }{ + {chatBID, c1ID}, + {chatB3ID, c3ID}, + {chatDID, c1ID}, + } { + execFixture( + `INSERT INTO chat_messages (chat_id, model_config_id, role, content, content_version) + VALUES ($1, $2, 'user', '[]'::jsonb, 2)`, + m.chatID, m.cfgID, + ) + } + + // Queued messages (FK-less) referencing configs via their chat's org. + execFixture( + `INSERT INTO chat_queued_messages (chat_id, model_config_id, content, created_by) + VALUES ($1, $2, '[]'::jsonb, $3)`, + chatBID, c1ID, user1ID, + ) + execFixture( + `INSERT INTO chat_queued_messages (chat_id, model_config_id, content, created_by) + VALUES ($1, $2, '[]'::jsonb, $3)`, + chatB3ID, c3ID, user2ID, + ) + execFixture( + `INSERT INTO chat_queued_messages (chat_id, model_config_id, content, created_by) + VALUES ($1, $2, '[]'::jsonb, $3)`, + chatDID, c1ID, user1ID, + ) + + // Debug runs (FK-less, attribution). + execFixture( + `INSERT INTO chat_debug_runs (id, chat_id, model_config_id, kind, status) + VALUES ($1, $2, $3, 'turn', 'finished')`, + uuid.New(), chatBID, c1ID, + ) + execFixture( + `INSERT INTO chat_debug_runs (id, chat_id, model_config_id, kind, status) + VALUES ($1, $2, $3, 'turn', 'finished')`, + uuid.New(), chatB3ID, c3ID, + ) + execFixture( + `INSERT INTO chat_debug_runs (id, chat_id, model_config_id, kind, status) + VALUES ($1, $2, $3, 'turn', 'finished')`, + uuid.New(), chatDID, c1ID, + ) + + // Compaction-threshold keys: c1 (live, users 1+2), c3 (deleted but + // referenced in orgB, user 1), c4 (deleted unreferenced, user 1: must + // never fan out), plus a non-threshold key that must stay untouched. + thresholdKey := func(id uuid.UUID) string { + return "chat_compaction_threshold_pct:" + id.String() + } + for _, tc := range []struct { + userID uuid.UUID + key string + value string + }{ + {user1ID, thresholdKey(c1ID), "80"}, + {user2ID, thresholdKey(c1ID), "75"}, + {user1ID, thresholdKey(c3ID), "60"}, + {user1ID, thresholdKey(c4ID), "55"}, + // Hostile keys: a malformed and an empty suffix. The up leaves + // them alone; the down must not abort on their uuid cast. + {user1ID, "chat_compaction_threshold_pct:not-a-uuid", "50"}, + {user1ID, "chat_compaction_threshold_pct:", "45"}, + {user1ID, "chat_personal_model_override:root", "chat_default"}, + } { + execFixture( + `INSERT INTO user_configs (user_id, key, value) VALUES ($1, $2, $3)`, + tc.userID, tc.key, tc.value, + ) + } + + // The fixtures were seeded post-schema in pre-explosion shape + // (default-org configs, references pointing at originals). The Stepper + // cannot stop mid-way, so the explosion of the seeded data is driven by + // rewinding the version row to the predecessor and stepping to latest, + // which re-applies this migration's up over the seeded rows. + _, err := sqlDB.ExecContext(ctx, fmt.Sprintf( + "TRUNCATE schema_migrations; INSERT INTO schema_migrations (version, dirty) VALUES (%d, false)", migrationVersion-1)) + require.NoError(t, err) + stepperUpToLatest(migrationVersion) + + // copyID resolves the copy of orig in org by natural attributes (the + // migration persists no mapping; this is also what the down relies on). + copyID := func(origID, orgID uuid.UUID) (uuid.UUID, bool) { + t.Helper() + var id uuid.UUID + err := sqlDB.QueryRowContext(ctx, + `SELECT cp.id FROM chat_model_configs cp + JOIN chat_model_configs orig ON orig.id = $1 + WHERE cp.organization_id = $2 + AND cp.model = orig.model + AND cp.ai_provider_id IS NOT DISTINCT FROM orig.ai_provider_id + AND cp.id <> orig.id`, origID, orgID).Scan(&id) + if err == sql.ErrNoRows { + return uuid.Nil, false + } + require.NoError(t, err) + return id, true + } + + // --- Per-org row counts --- + // Default org keeps its 6 originals; orgB gets 4 live copies (c1, c2, + // c5, late) plus the referenced-deleted c3 copy; orgC (zero-member) gets + // the full live set (4) and nothing deleted; orgD gets nothing. + assertCount := func(orgID uuid.UUID, want int, msg string) { + t.Helper() + var got int + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT COUNT(*) FROM chat_model_configs WHERE organization_id = $1", orgID).Scan(&got)) + require.Equal(t, want, got, msg) + } + assertCount(defaultOrgID, 6, "default org keeps only its originals") + assertCount(orgBID, 5, "orgB: 4 live fan-out + referenced-deleted c3") + assertCount(orgCID, 4, "orgC: full live set, no deleted copies") + assertCount(orgDID, 0, "soft-deleted org receives no copies") + + var totalConfigs int + require.NoError(t, sqlDB.QueryRowContext(ctx, "SELECT COUNT(*) FROM chat_model_configs").Scan(&totalConfigs)) + require.Equal(t, 15, totalConfigs, "up: 6 originals and 9 copies") + + var duplicateConfigs int + require.NoError(t, sqlDB.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM ( + SELECT organization_id, ai_provider_id, model + FROM chat_model_configs + GROUP BY organization_id, ai_provider_id, model + HAVING COUNT(*) <> 1 + ) duplicates + `).Scan(&duplicateConfigs)) + require.Zero(t, duplicateConfigs, "each organization has one config per provider and model") + + // Copies preserve deleted state and deletion timestamps. + c3CopyB, ok := copyID(c3ID, orgBID) + require.True(t, ok, "orgB received the referenced-deleted c3 copy") + var mismatchedDeletedState int + require.NoError(t, sqlDB.QueryRowContext(ctx, ` + SELECT COUNT(*) + FROM chat_model_configs cp + JOIN chat_model_configs orig + ON orig.organization_id = $1 + AND orig.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id + AND orig.model = cp.model + WHERE cp.organization_id IN ($2, $3) + AND (cp.deleted IS DISTINCT FROM orig.deleted + OR cp.deleted_at IS DISTINCT FROM orig.deleted_at) + `, defaultOrgID, orgBID, orgCID).Scan(&mismatchedDeletedState)) + require.Zero(t, mismatchedDeletedState) + + // c4 is copied nowhere. + for _, orgID := range []uuid.UUID{orgBID, orgCID, orgDID} { + _, ok := copyID(c4ID, orgID) + require.False(t, ok, "unreferenced deleted c4 must not be copied") + } + + // --- Exactly one live default per org that received copies --- + for _, orgID := range []uuid.UUID{defaultOrgID, orgBID, orgCID} { + var defaults int + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT COUNT(*) FROM chat_model_configs WHERE organization_id = $1 AND is_default AND NOT deleted", orgID).Scan(&defaults)) + require.Equal(t, 1, defaults, "exactly one live default per org") + } + + // --- Remaps --- + // Chats in live orgB remap to same-org copies; the chat in soft-deleted + // orgD keeps the original reference. + c1CopyB, ok := copyID(c1ID, orgBID) + require.True(t, ok) + assertChatPinned := func(chatID, want uuid.UUID, msg string) { + t.Helper() + var got uuid.UUID + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT last_model_config_id FROM chats WHERE id = $1", chatID).Scan(&got)) + require.Equal(t, want, got, msg) + } + assertChatPinned(chatBID, c1CopyB, "orgB chat remapped to same-org live copy") + assertChatPinned(chatB3ID, c3CopyB, "orgB chat on deleted model remapped to the deleted copy") + assertChatPinned(chatDID, c1ID, "soft-deleted org chat keeps original reference") + + var msgCfg uuid.UUID + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatBID).Scan(&msgCfg)) + require.Equal(t, c1CopyB, msgCfg, "orgB message remapped") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) + require.Equal(t, c3CopyB, msgCfg, "orgB deleted-model message remapped") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatDID).Scan(&msgCfg)) + require.Equal(t, c1ID, msgCfg, "orgD message untouched") + + // Queued messages (FK-less): orgB remapped, orgD untouched. + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatBID).Scan(&msgCfg)) + require.Equal(t, c1CopyB, msgCfg, "orgB queued message remapped") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) + require.Equal(t, c3CopyB, msgCfg, "orgB deleted-model queued message remapped") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatDID).Scan(&msgCfg)) + require.Equal(t, c1ID, msgCfg, "orgD queued message untouched") + + // Debug runs (FK-less): orgB remapped, orgD untouched. + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatBID).Scan(&msgCfg)) + require.Equal(t, c1CopyB, msgCfg, "orgB debug run remapped") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) + require.Equal(t, c3CopyB, msgCfg, "orgB deleted-model debug run remapped") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatDID).Scan(&msgCfg)) + require.Equal(t, c1ID, msgCfg, "orgD debug run untouched") + + // No reference in a live non-default org still points at a default-org + // config. + var dangling int + require.NoError(t, sqlDB.QueryRowContext(ctx, + `SELECT COUNT(*) FROM chats c + JOIN chat_model_configs cmc ON cmc.id = c.last_model_config_id + JOIN organizations def ON def.id = cmc.organization_id AND def.is_default + JOIN organizations co ON co.id = c.organization_id + WHERE NOT co.is_default AND NOT co.deleted`).Scan(&dangling)) + require.Zero(t, dangling, "no live-org chat references a default-org config") + + // --- ACL re-key and backfill --- + // Each copy carries only its target organization's everyone entry. + for _, orgID := range []uuid.UUID{orgBID, orgCID} { + rows, err := sqlDB.QueryContext(ctx, + "SELECT group_acl FROM chat_model_configs WHERE organization_id = $1", orgID) + require.NoError(t, err) + for rows.Next() { + var groupACL []byte + require.NoError(t, rows.Scan(&groupACL)) + require.JSONEq(t, string(mustJSON(t, map[string]any{ + orgID.String(): map[string]any{"permissions": []string{"read"}}, + })), string(groupACL)) + } + require.NoError(t, rows.Err()) + require.NoError(t, rows.Close()) + } + // The pre-existing row with an empty group_acl was backfilled. + var aclRaw []byte + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT group_acl FROM chat_model_configs WHERE id = $1", aclLateConfigID).Scan(&aclRaw)) + require.JSONEq(t, string(mustJSON(t, map[string]any{ + defaultOrgID.String(): map[string]any{"permissions": []string{"read"}}, + })), string(aclRaw), "row missing the everyone entry is backfilled") + + // --- Threshold fan-out --- + // c1 keys fanned out to the orgB and orgC copies for both users (follows + // copies, not membership); the c3 key fanned out to the orgB copy only; + // c4 produced nothing. + c1CopyC, ok := copyID(c1ID, orgCID) + require.True(t, ok) + for _, tc := range []struct { + userID uuid.UUID + cfgID uuid.UUID + value string + }{ + {user1ID, c1CopyB, "80"}, + {user1ID, c1CopyC, "80"}, + {user2ID, c1CopyB, "75"}, + {user2ID, c1CopyC, "75"}, + {user1ID, c3CopyB, "60"}, + } { + var value string + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT value FROM user_configs WHERE user_id = $1 AND key = $2", + tc.userID, thresholdKey(tc.cfgID)).Scan(&value)) + require.Equal(t, tc.value, value, "threshold value copied to config %s for user %s", tc.cfgID, tc.userID) + } + var validThresholdCount int + require.NoError(t, sqlDB.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM user_configs + WHERE key ~ '^chat_compaction_threshold_pct:[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' + `).Scan(&validThresholdCount)) + require.Equal(t, 9, validThresholdCount, "4 original and 5 fanned-out threshold keys") + // c3 produced no orgC key (its only copy is in orgB). + var c3CCount int + require.NoError(t, sqlDB.QueryRowContext(ctx, + `SELECT COUNT(*) FROM user_configs uc + JOIN chat_model_configs cp ON cp.id = substring(uc.key FROM 'chat_compaction_threshold_pct:(.*)')::uuid + WHERE uc.key LIKE 'chat_compaction_threshold_pct:%' + AND substring(uc.key FROM 'chat_compaction_threshold_pct:(.*)') ~ '^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' + AND cp.organization_id = $1 AND cp.model = 'gpt-4-legacy'`, + orgCID).Scan(&c3CCount)) + require.Zero(t, c3CCount, "deleted config with no orgC copy produces no orgC key") + // c4 (deleted, unreferenced): the only ancient-model threshold key is the + // seeded original. + var c4Keys []string + rows, err := sqlDB.QueryContext(ctx, + `SELECT key FROM user_configs WHERE key LIKE 'chat_compaction_threshold_pct:%' + AND substring(key FROM 'chat_compaction_threshold_pct:(.*)') ~ '^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' + AND substring(key FROM 'chat_compaction_threshold_pct:(.*)')::uuid = $1`, c4ID) + require.NoError(t, err) + for rows.Next() { + var k string + require.NoError(t, rows.Scan(&k)) + c4Keys = append(c4Keys, k) + } + require.NoError(t, rows.Err()) + require.NoError(t, rows.Close()) + require.Equal(t, []string{thresholdKey(c4ID)}, c4Keys, "unreferenced deleted c4 fans out zero keys") + // Seeded original keys survive; the non-threshold key is untouched. + for _, key := range []string{thresholdKey(c1ID), thresholdKey(c3ID)} { + var exists bool + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT EXISTS(SELECT 1 FROM user_configs WHERE user_id = $1 AND key = $2)", + user1ID, key).Scan(&exists)) + require.True(t, exists, "original threshold key survives") + } + var overrideValue string + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT value FROM user_configs WHERE user_id = $1 AND key = 'chat_personal_model_override:root'", + user1ID).Scan(&overrideValue)) + require.Equal(t, "chat_default", overrideValue, "non-threshold keys are untouched") + + // --- Down round-trip on the exploded state. The framework's Stepper + // cannot commit after exactly one down step (its driver commits only + // when the stepper exhausts), so the down file is executed directly: + // migrations are plain SQL and this is the same text golang-migrate + // would run. + downSQL, err := fs.ReadFile(migrations.MigrationFS(), "000566_chat_model_config_org_explosion.down.sql") + require.NoError(t, err) + _, err = sqlDB.ExecContext(ctx, string(downSQL)) + require.NoError(t, err) + + // Copies are gone; references restored to the default-org originals. + assertCount(defaultOrgID, 6, "down: default org keeps its originals") + assertCount(orgBID, 0, "down: copies deleted from orgB") + assertCount(orgCID, 0, "down: copies deleted from orgC") + require.NoError(t, sqlDB.QueryRowContext(ctx, "SELECT COUNT(*) FROM chat_model_configs").Scan(&totalConfigs)) + require.Equal(t, 6, totalConfigs, "down: only default-org originals remain") + assertChatPinned(chatBID, c1ID, "down: orgB chat restored to original c1") + assertChatPinned(chatB3ID, c3ID, "down: orgB chat restored to original c3") + assertChatPinned(chatDID, c1ID, "down: orgD chat unchanged") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatBID).Scan(&msgCfg)) + require.Equal(t, c1ID, msgCfg, "down: orgB message restored") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) + require.Equal(t, c3ID, msgCfg, "down: orgB deleted-model message restored") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatBID).Scan(&msgCfg)) + require.Equal(t, c1ID, msgCfg, "down: orgB queued message restored") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) + require.Equal(t, c3ID, msgCfg, "down: orgB deleted-model queued message restored") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatBID).Scan(&msgCfg)) + require.Equal(t, c1ID, msgCfg, "down: orgB debug run restored") + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) + require.Equal(t, c3ID, msgCfg, "down: orgB deleted-model debug run restored") + + // Fanned-out threshold keys are gone; exactly the seeded valid key set + // remains. The two hostile keys (malformed/empty suffix) are pruned as + // dangling: they cannot name an existing config, and the down must not + // abort on their uuid cast. + var thresholdCount int + require.NoError(t, sqlDB.QueryRowContext(ctx, + "SELECT COUNT(*) FROM user_configs WHERE key LIKE 'chat_compaction_threshold_pct:%'").Scan(&thresholdCount)) + require.Equal(t, 4, thresholdCount, "down: only the 4 seeded valid threshold keys remain") + + // Step forward again: the down executed outside the driver, so rewind + // the version row and re-apply the up. Shape re-asserted at count + // level. + _, err = sqlDB.ExecContext(ctx, fmt.Sprintf( + "TRUNCATE schema_migrations; INSERT INTO schema_migrations (version, dirty) VALUES (%d, false)", migrationVersion-1)) + require.NoError(t, err) + stepperUpToLatest(migrationVersion) + assertCount(defaultOrgID, 6, "re-up: default org keeps its originals") + assertCount(orgBID, 5, "re-up: orgB copies recreated") + assertCount(orgCID, 4, "re-up: orgC copies recreated") + assertCount(orgDID, 0, "re-up: soft-deleted org receives no copies") + require.NoError(t, sqlDB.QueryRowContext(ctx, "SELECT COUNT(*) FROM chat_model_configs").Scan(&totalConfigs)) + require.Equal(t, 15, totalConfigs, "re-up: 6 originals and 9 copies") + require.NoError(t, sqlDB.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM ( + SELECT organization_id, ai_provider_id, model + FROM chat_model_configs + GROUP BY organization_id, ai_provider_id, model + HAVING COUNT(*) <> 1 + ) duplicates + `).Scan(&duplicateConfigs)) + require.Zero(t, duplicateConfigs, "re-up: each organization has one config per provider and model") + require.NoError(t, sqlDB.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM user_configs + WHERE key ~ '^chat_compaction_threshold_pct:[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' + `).Scan(&validThresholdCount)) + require.Equal(t, 9, validThresholdCount, "re-up: threshold keys fan out once") + assertChatPinned(chatDID, c1ID, "re-up: orgD still untouched") +} diff --git a/coderd/database/migrations/testdata/fixtures/000566_chat_model_config_org_explosion.up.sql b/coderd/database/migrations/testdata/fixtures/000566_chat_model_config_org_explosion.up.sql new file mode 100644 index 00000000000..ae2b62dd6c5 --- /dev/null +++ b/coderd/database/migrations/testdata/fixtures/000566_chat_model_config_org_explosion.up.sql @@ -0,0 +1,33 @@ +-- Fixture for 000566 (org explosion cutover). Fixtures apply right after +-- their migration runs, so this executes after the explosion has copied +-- the default organization's configs into every live organization. It +-- seeds one organically created config in the non-default organization +-- from fixture 000291, so later migrations run over configs that the +-- explosion did not create. +INSERT INTO chat_model_configs ( + id, + model, + display_name, + enabled, + is_default, + context_limit, + compression_threshold, + ai_provider_id, + organization_id, + group_acl, + created_at, + updated_at +) VALUES ( + '566c0001-0000-4000-8000-000000000001', + 'gpt-5.2-org-fixture', + 'Fixture Org Model 566', + TRUE, + FALSE, + 128000, + 70, + 'a52c6f0e-7d4b-4e1a-9c3f-2b8d5e6f7a8b', + '20362772-802a-4a72-8e4f-3648b4bfd168', + jsonb_build_object('20362772-802a-4a72-8e4f-3648b4bfd168', jsonb_build_object('permissions', jsonb_build_array('read'))), + '2024-01-01 00:00:00+00', + '2024-01-01 00:00:00+00' +); diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 47597723a7f..b5a7e8a4bd0 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -8493,19 +8493,20 @@ func (q *sqlQuerier) GetChatMessagesForPromptByChatID(ctx context.Context, chatI } const getChatModelConfigsForTelemetry = `-- name: GetChatModelConfigsForTelemetry :many -SELECT cmc.id, ap.type::text AS provider, cmc.model, cmc.context_limit, cmc.enabled, cmc.is_default +SELECT cmc.id, ap.type::text AS provider, cmc.model, cmc.context_limit, cmc.enabled, cmc.is_default, cmc.organization_id FROM chat_model_configs cmc JOIN ai_providers ap ON ap.id = cmc.ai_provider_id WHERE cmc.deleted = false ` type GetChatModelConfigsForTelemetryRow struct { - ID uuid.UUID `db:"id" json:"id"` - Provider string `db:"provider" json:"provider"` - Model string `db:"model" json:"model"` - ContextLimit int64 `db:"context_limit" json:"context_limit"` - Enabled bool `db:"enabled" json:"enabled"` - IsDefault bool `db:"is_default" json:"is_default"` + ID uuid.UUID `db:"id" json:"id"` + Provider string `db:"provider" json:"provider"` + Model string `db:"model" json:"model"` + ContextLimit int64 `db:"context_limit" json:"context_limit"` + Enabled bool `db:"enabled" json:"enabled"` + IsDefault bool `db:"is_default" json:"is_default"` + OrganizationID uuid.UUID `db:"organization_id" json:"organization_id"` } // Returns all model configurations for telemetry snapshot collection. @@ -8526,6 +8527,7 @@ func (q *sqlQuerier) GetChatModelConfigsForTelemetry(ctx context.Context) ([]Get &i.ContextLimit, &i.Enabled, &i.IsDefault, + &i.OrganizationID, ); err != nil { return nil, err } diff --git a/coderd/database/queries/chats.sql b/coderd/database/queries/chats.sql index 290e0c17f9d..5f3d9aa4cac 100644 --- a/coderd/database/queries/chats.sql +++ b/coderd/database/queries/chats.sql @@ -2291,7 +2291,7 @@ GROUP BY cm.chat_id; -- name: GetChatModelConfigsForTelemetry :many -- Returns all model configurations for telemetry snapshot collection. -- deleted = false guarantees ai_provider_id is non-null, so INNER JOIN is safe. -SELECT cmc.id, ap.type::text AS provider, cmc.model, cmc.context_limit, cmc.enabled, cmc.is_default +SELECT cmc.id, ap.type::text AS provider, cmc.model, cmc.context_limit, cmc.enabled, cmc.is_default, cmc.organization_id FROM chat_model_configs cmc JOIN ai_providers ap ON ap.id = cmc.ai_provider_id WHERE cmc.deleted = false; diff --git a/coderd/telemetry/telemetry.go b/coderd/telemetry/telemetry.go index 28e69a248b3..f35873a98c2 100644 --- a/coderd/telemetry/telemetry.go +++ b/coderd/telemetry/telemetry.go @@ -2318,12 +2318,13 @@ func ConvertChatMessageSummary(dbRow database.GetChatMessageSummariesPerChatRow) // telemetry ChatModelConfig. func ConvertChatModelConfig(dbRow database.GetChatModelConfigsForTelemetryRow) ChatModelConfig { return ChatModelConfig{ - ID: dbRow.ID, - Provider: dbRow.Provider, - Model: dbRow.Model, - ContextLimit: dbRow.ContextLimit, - Enabled: dbRow.Enabled, - IsDefault: dbRow.IsDefault, + ID: dbRow.ID, + OrganizationID: dbRow.OrganizationID, + Provider: dbRow.Provider, + Model: dbRow.Model, + ContextLimit: dbRow.ContextLimit, + Enabled: dbRow.Enabled, + IsDefault: dbRow.IsDefault, } } @@ -2611,12 +2612,13 @@ type ChatMessageSummary struct { // ChatModelConfig contains model configuration metadata for // telemetry. Sensitive fields like API keys are excluded. type ChatModelConfig struct { - ID uuid.UUID `json:"id"` - Provider string `json:"provider"` - Model string `json:"model"` - ContextLimit int64 `json:"context_limit"` - Enabled bool `json:"enabled"` - IsDefault bool `json:"is_default"` + ID uuid.UUID `json:"id"` + OrganizationID uuid.UUID `json:"organization_id"` + Provider string `json:"provider"` + Model string `json:"model"` + ContextLimit int64 `json:"context_limit"` + Enabled bool `json:"enabled"` + IsDefault bool `json:"is_default"` } // ChatDiffStatusSummary contains aggregate PR counts across all diff --git a/coderd/telemetry/telemetry_test.go b/coderd/telemetry/telemetry_test.go index 875882388f4..8be5f66559c 100644 --- a/coderd/telemetry/telemetry_test.go +++ b/coderd/telemetry/telemetry_test.go @@ -1935,6 +1935,7 @@ func TestChatsTelemetry(t *testing.T) { cfg1, ok := configMap[modelCfg.ID] require.True(t, ok) + assert.Equal(t, org.ID, cfg1.OrganizationID) assert.Equal(t, "anthropic", cfg1.Provider) assert.Equal(t, "claude-sonnet-4-20250514", cfg1.Model) assert.Equal(t, int64(200000), cfg1.ContextLimit) @@ -1943,6 +1944,7 @@ func TestChatsTelemetry(t *testing.T) { cfg2, ok := configMap[modelCfg2.ID] require.True(t, ok) + assert.Equal(t, org.ID, cfg2.OrganizationID) assert.Equal(t, "openai", cfg2.Provider) assert.Equal(t, "gpt-4o", cfg2.Model) assert.Equal(t, int64(128000), cfg2.ContextLimit) From a4adf6008ad157c5567839a37caa64bbe342fdb3 Mon Sep 17 00:00:00 2001 From: Ethan Dickson Date: Mon, 10 Aug 2026 09:12:00 +0000 Subject: [PATCH 2/3] fix: enforce organization-local chat model configs --- ...6_chat_model_config_org_explosion.down.sql | 142 +++-------- ...566_chat_model_config_org_explosion.up.sql | 172 ++++--------- coderd/database/migrations/migrate.go | 8 - coderd/database/migrations/migrate_test.go | 85 ++----- coderd/exp_chats.go | 48 ++-- coderd/exp_chats_test.go | 181 ++++++++++++- coderd/x/chatd/chatd.go | 119 ++------- coderd/x/chatd/chatd_internal_test.go | 65 ++--- coderd/x/chatd/compaction_override.go | 3 + .../compaction_override_internal_test.go | 28 +++ coderd/x/chatd/configcache.go | 13 +- coderd/x/chatd/configcache_internal_test.go | 71 +----- coderd/x/chatd/subagent.go | 12 +- coderd/x/chatd/subagent_internal_test.go | 237 +++--------------- coderd/x/chatd/title_override.go | 2 +- .../x/chatd/title_override_internal_test.go | 61 +++-- enterprise/coderd/exp_chats_test.go | 19 +- 17 files changed, 519 insertions(+), 747 deletions(-) diff --git a/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql b/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql index 27e3f453c4e..da93da6ec57 100644 --- a/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql +++ b/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql @@ -1,113 +1,53 @@ --- DOWN for the chat model config org explosion. UNSUPPORTED and best-effort --- per operator ruling: there is no persisted provenance, so this down --- cannot distinguish a copy from an organically created non-default-org row --- that happens to share (ai_provider_id, model) with a default-org row. It --- must run green and must never lose chats; fidelity loss on pathological --- duplicates is accepted. --- --- Copy identification: a non-default-org row is treated as a copy iff a --- default-org row exists with the same (ai_provider_id, model). Retarget --- resolves that default-org row deterministically with DISTINCT ON ordered --- by (created_at ASC, id ASC) so duplicates pick one stable row. +-- This migration is best-effort because copied configs have no persisted +-- provenance. A non-default config is treated as a copy when a default-org +-- config has the same provider and model. +CREATE TEMPORARY TABLE model_config_copy_map ( + copy_id uuid PRIMARY KEY, + orig_id uuid NOT NULL +) ON COMMIT DROP; --- Restore chats.last_model_config_id from copies back to the default-org --- original matched by (ai_provider_id, model). -UPDATE chats c -SET last_model_config_id = orig.id +INSERT INTO model_config_copy_map (copy_id, orig_id) +SELECT cp.id, orig.id FROM chat_model_configs cp +JOIN organizations copy_org + ON copy_org.id = cp.organization_id + AND NOT copy_org.is_default JOIN LATERAL ( - SELECT d.id - FROM chat_model_configs d - JOIN organizations def ON def.id = d.organization_id AND def.is_default - WHERE d.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id - AND d.model = cp.model - ORDER BY d.created_at ASC, d.id ASC + SELECT default_config.id + FROM chat_model_configs default_config + JOIN organizations default_org + ON default_org.id = default_config.organization_id + AND default_org.is_default + WHERE default_config.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id + AND default_config.model = cp.model + ORDER BY default_config.created_at ASC, default_config.id ASC LIMIT 1 -) orig ON true -WHERE c.last_model_config_id = cp.id - AND NOT EXISTS (SELECT 1 FROM organizations odef - WHERE odef.id = cp.organization_id AND odef.is_default); +) orig ON true; + +UPDATE chats c +SET last_model_config_id = m.orig_id +FROM model_config_copy_map m +WHERE c.last_model_config_id = m.copy_id; --- Restore chat_messages.model_config_id from copies back to originals. UPDATE chat_messages mm -SET model_config_id = orig.id -FROM chat_model_configs cp -JOIN LATERAL ( - SELECT d.id - FROM chat_model_configs d - JOIN organizations def ON def.id = d.organization_id AND def.is_default - WHERE d.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id - AND d.model = cp.model - ORDER BY d.created_at ASC, d.id ASC - LIMIT 1 -) orig ON true -WHERE mm.model_config_id = cp.id - AND NOT EXISTS (SELECT 1 FROM organizations odef - WHERE odef.id = cp.organization_id AND odef.is_default); +SET model_config_id = m.orig_id +FROM model_config_copy_map m +WHERE mm.model_config_id = m.copy_id; --- Restore chat_queued_messages.model_config_id from copies back to originals. UPDATE chat_queued_messages q -SET model_config_id = orig.id -FROM chat_model_configs cp -JOIN LATERAL ( - SELECT d.id - FROM chat_model_configs d - JOIN organizations def ON def.id = d.organization_id AND def.is_default - WHERE d.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id - AND d.model = cp.model - ORDER BY d.created_at ASC, d.id ASC - LIMIT 1 -) orig ON true -WHERE q.model_config_id = cp.id - AND NOT EXISTS (SELECT 1 FROM organizations odef - WHERE odef.id = cp.organization_id AND odef.is_default); +SET model_config_id = m.orig_id +FROM model_config_copy_map m +WHERE q.model_config_id = m.copy_id; --- Restore chat_debug_runs.model_config_id from copies back to originals. UPDATE chat_debug_runs d -SET model_config_id = orig.id -FROM chat_model_configs cp -JOIN LATERAL ( - SELECT dc.id - FROM chat_model_configs dc - JOIN organizations def ON def.id = dc.organization_id AND def.is_default - WHERE dc.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id - AND dc.model = cp.model - ORDER BY dc.created_at ASC, dc.id ASC - LIMIT 1 -) orig ON true -WHERE d.model_config_id = cp.id - AND NOT EXISTS (SELECT 1 FROM organizations odef - WHERE odef.id = cp.organization_id AND odef.is_default); - --- Delete copied chat_model_configs (non-default-org rows whose --- (ai_provider_id, model) matches a default-org row). References were --- retargeted above, so the deletes cannot violate the chats/chat_messages --- FKs. -DELETE FROM chat_model_configs cp -WHERE EXISTS ( - SELECT 1 FROM chat_model_configs orig - JOIN organizations def ON def.id = orig.organization_id AND def.is_default - WHERE orig.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id - AND orig.model = cp.model -) -AND NOT EXISTS ( - SELECT 1 FROM organizations odef - WHERE odef.id = cp.organization_id AND odef.is_default -); +SET model_config_id = m.orig_id +FROM model_config_copy_map m +WHERE d.model_config_id = m.copy_id; --- Best-effort threshold-key cleanup: delete compaction-threshold keys whose --- embedded config id no longer exists anywhere after the copy deletes. --- Original keys survive because default-org originals always survive. --- Keys with a malformed or empty suffix are guarded BEFORE the uuid cast --- (they cannot name an existing config, so they are pruned like any other --- dangling key) instead of aborting the down. DELETE FROM user_configs uc -WHERE uc.key LIKE 'chat_compaction_threshold_pct:%' - AND NOT EXISTS ( - SELECT 1 FROM chat_model_configs cmc - WHERE cmc.id = ( - SELECT substring(uc.key FROM 'chat_compaction_threshold_pct:(.*)')::uuid - WHERE substring(uc.key FROM 'chat_compaction_threshold_pct:(.*)') - ~ '^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' - ) - ); +USING model_config_copy_map m +WHERE uc.key = 'chat_compaction_threshold_pct:' || m.copy_id::text; + +DELETE FROM chat_model_configs cmc +USING model_config_copy_map m +WHERE cmc.id = m.copy_id; diff --git a/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql b/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql index 6476d72047a..ee0ed461d69 100644 --- a/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql +++ b/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql @@ -1,17 +1,6 @@ --- Explode default-org chat model configs to every live non-default --- organization (CODAGT-709, stage 3 of 3: org-scoping cutover). After this --- migration every live org owns a full set of model configs and all --- references inside live non-default orgs point at same-org rows. --- --- Mapping design (operator ruling): NO provenance column, NO persisted --- mapping of any kind. A transaction-scoped TEMPORARY lookup table --- (orig_id, org_id, copy_id) ON COMMIT DROP maps each default-org original --- to its per-org copy; the copy insert, all four reference remaps, and the --- compaction-threshold fan-out join it. The table vanishes when the --- migration framework commits this migration's transaction, so nothing --- mapping-related persists. Copy ids come from gen_random_uuid(); no --- hash-derived ids (md5() errors on FIPS-mode PostgreSQL builds). - +-- Copy default-organization chat model configs into each live non-default +-- organization. Referenced soft-deleted configs are copied only into the +-- organizations that reference them. CREATE TEMPORARY TABLE model_config_copy_map ( orig_id uuid NOT NULL, org_id uuid NOT NULL, @@ -19,76 +8,46 @@ CREATE TEMPORARY TABLE model_config_copy_map ( PRIMARY KEY (orig_id, org_id) ) ON COMMIT DROP; --- (a) Stage LIVE default-org chat_model_configs x every live non-default org --- in the temp map with a fresh copy id, then insert the copies. Staging the --- id in the map first lets the remap statements below resolve copies without --- recomputing anything. INSERT INTO model_config_copy_map (orig_id, org_id, copy_id) SELECT cmc.id, o.id, gen_random_uuid() FROM chat_model_configs cmc JOIN organizations def ON def.id = cmc.organization_id AND def.is_default CROSS JOIN organizations o -WHERE NOT o.is_default AND NOT o.deleted - AND NOT cmc.deleted; - --- Copies inherit every behavioral field from the original, including --- created_at/updated_at and created_by/updated_by: a copy is the same --- logical config re-homed, and the audit-facing identity of who configured --- it survives the explosion. group_acl is re-keyed to the copy's org (the --- Everyone group of an organization always has the organization's own ID, --- see 000058) carrying the original's entry verbatim, so members of the --- target org keep read access through the everyone entry. -INSERT INTO chat_model_configs - (id, model, display_name, created_by, updated_by, enabled, is_default, - deleted, deleted_at, created_at, updated_at, context_limit, - compression_threshold, options, ai_provider_id, organization_id, - group_acl, user_acl) -SELECT - m.copy_id, - cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, - cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, - cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, - cmc.ai_provider_id, m.org_id, - jsonb_build_object( - m.org_id::text, - COALESCE(cmc.group_acl -> cmc.organization_id::text, - '{"permissions": ["read"]}'::jsonb) - ), - '{}'::jsonb -FROM model_config_copy_map m -JOIN chat_model_configs cmc ON cmc.id = m.orig_id -WHERE NOT cmc.deleted; - --- (a2) Stage + copy SOFT-DELETED default-org chat_model_configs ONLY to live --- non-default orgs that actually reference them. A reference is any of: --- chats.last_model_config_id, chat_messages.model_config_id (via chat), --- chat_queued_messages.model_config_id (via chat), or --- chat_debug_runs.model_config_id (via chat) pointing at the deleted config. --- Copies keep deleted/deleted_at so every historical reference has an --- FK-valid, attribution-preserving target without resurrecting the config. -INSERT INTO model_config_copy_map (orig_id, org_id, copy_id) -SELECT DISTINCT ON (cmc.id, o.id) cmc.id, o.id, gen_random_uuid() -FROM chat_model_configs cmc -JOIN organizations def ON def.id = cmc.organization_id AND def.is_default -JOIN organizations o ON NOT o.is_default AND NOT o.deleted -WHERE cmc.deleted +WHERE NOT o.is_default + AND NOT o.deleted AND ( - EXISTS (SELECT 1 FROM chats c - WHERE c.last_model_config_id = cmc.id AND c.organization_id = o.id) - OR - EXISTS (SELECT 1 FROM chat_messages mm - JOIN chats c ON c.id = mm.chat_id - WHERE mm.model_config_id = cmc.id AND c.organization_id = o.id) - OR - EXISTS (SELECT 1 FROM chat_queued_messages q - JOIN chats c ON c.id = q.chat_id - WHERE q.model_config_id = cmc.id AND c.organization_id = o.id) - OR - EXISTS (SELECT 1 FROM chat_debug_runs d - JOIN chats c ON c.id = d.chat_id - WHERE d.model_config_id = cmc.id AND c.organization_id = o.id) + NOT cmc.deleted + OR EXISTS ( + SELECT 1 + FROM chats c + WHERE c.last_model_config_id = cmc.id + AND c.organization_id = o.id + ) + OR EXISTS ( + SELECT 1 + FROM chat_messages mm + JOIN chats c ON c.id = mm.chat_id + WHERE mm.model_config_id = cmc.id + AND c.organization_id = o.id + ) + OR EXISTS ( + SELECT 1 + FROM chat_queued_messages q + JOIN chats c ON c.id = q.chat_id + WHERE q.model_config_id = cmc.id + AND c.organization_id = o.id + ) + OR EXISTS ( + SELECT 1 + FROM chat_debug_runs d + JOIN chats c ON c.id = d.chat_id + WHERE d.model_config_id = cmc.id + AND c.organization_id = o.id + ) ); +-- Each copy retains the original behavior and audit fields. The everyone +-- group ACL is re-keyed to the destination organization. INSERT INTO chat_model_configs (id, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, @@ -96,30 +55,38 @@ INSERT INTO chat_model_configs group_acl, user_acl) SELECT m.copy_id, - cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, - cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, - cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, - cmc.ai_provider_id, m.org_id, + cmc.model, + cmc.display_name, + cmc.created_by, + cmc.updated_by, + cmc.enabled, + cmc.is_default, + cmc.deleted, + cmc.deleted_at, + cmc.created_at, + cmc.updated_at, + cmc.context_limit, + cmc.compression_threshold, + cmc.options, + cmc.ai_provider_id, + m.org_id, jsonb_build_object( m.org_id::text, - COALESCE(cmc.group_acl -> cmc.organization_id::text, - '{"permissions": ["read"]}'::jsonb) + COALESCE( + cmc.group_acl -> cmc.organization_id::text, + '{"permissions": ["read"]}'::jsonb + ) ), '{}'::jsonb FROM model_config_copy_map m -JOIN chat_model_configs cmc ON cmc.id = m.orig_id -WHERE cmc.deleted; +JOIN chat_model_configs cmc ON cmc.id = m.orig_id; --- (b) Remap chats.last_model_config_id in live non-default orgs to the --- same-org copy via the temp map. Soft-deleted orgs have no map rows, so --- their chats keep original references. UPDATE chats c SET last_model_config_id = m.copy_id FROM model_config_copy_map m WHERE c.last_model_config_id = m.orig_id AND m.org_id = c.organization_id; --- (b2) Remap chat_messages.model_config_id via the owning chat's org. UPDATE chat_messages mm SET model_config_id = m.copy_id FROM chats c, model_config_copy_map m @@ -127,9 +94,6 @@ WHERE c.id = mm.chat_id AND mm.model_config_id = m.orig_id AND m.org_id = c.organization_id; --- (b3) Remap chat_queued_messages.model_config_id via the owning chat's org. --- The column has no FK, so dangling ids would not fail. The remap keeps a --- queued message's promoted model inside its chat's org. UPDATE chat_queued_messages q SET model_config_id = m.copy_id FROM chats c, model_config_copy_map m @@ -137,8 +101,6 @@ WHERE c.id = q.chat_id AND q.model_config_id = m.orig_id AND m.org_id = c.organization_id; --- (b4) Remap chat_debug_runs.model_config_id via the owning chat's org. --- The column is FK-less and stores attribution only. UPDATE chat_debug_runs d SET model_config_id = m.copy_id FROM chats c, model_config_copy_map m @@ -146,33 +108,9 @@ WHERE c.id = d.chat_id AND d.model_config_id = m.orig_id AND m.org_id = c.organization_id; --- (c) Fan out user_configs compaction-threshold keys. A key --- 'chat_compaction_threshold_pct:' earns one row per copy of that --- original in the temp map, same user, same value, key rewritten to the --- copy id. The fan-out is copy-precise by construction (it can only --- produce keys for copies that exist): live originals reach every live --- org, soft-deleted originals reach only the orgs that received a --- referenced copy, and an original with zero map rows (deleted and --- unreferenced) produces nothing. Original keys stay: they reference --- default-org originals, still valid. The fan-out is deliberately NOT --- membership-filtered: chats pinned to deleted models are the norm, and a --- threshold must keep resolving for any chat that lands on a copy. The PK --- (user_id, key) cannot collide because copy ids are fresh and no existing --- key embeds a copy id; ON CONFLICT DO NOTHING is belt-and-braces only. INSERT INTO user_configs (user_id, key, value) SELECT uc.user_id, 'chat_compaction_threshold_pct:' || m.copy_id::text, uc.value FROM user_configs uc JOIN model_config_copy_map m ON uc.key = 'chat_compaction_threshold_pct:' || m.orig_id::text ON CONFLICT (user_id, key) DO NOTHING; - --- (d) Seed the everyone-in-org read entry on any existing row whose --- group_acl lacks its own org's key. This covers rows written by older --- binaries during a rolling upgrade. The entry's permissions are preserved --- when an entry already exists for another org's key shape. -UPDATE chat_model_configs -SET group_acl = jsonb_build_object( - organization_id::text, - jsonb_build_object('permissions', jsonb_build_array('read'::text)) -) || group_acl -WHERE NOT (group_acl ? organization_id::text); diff --git a/coderd/database/migrations/migrate.go b/coderd/database/migrations/migrate.go index ec03a606f66..50a931c902f 100644 --- a/coderd/database/migrations/migrate.go +++ b/coderd/database/migrations/migrate.go @@ -23,14 +23,6 @@ import ( //go:embed *.sql var migrations embed.FS -// MigrationFS exposes the embedded migration files, for tests that need to -// execute a single migration's SQL outside the migrate driver (the driver's -// transaction commits only when a stepper exhausts, which mid-test -// down-then-up cycles cannot wait for). -func MigrationFS() fs.FS { - return migrations -} - var ( migrationsHash string migrationsHashOnce sync.Once diff --git a/coderd/database/migrations/migrate_test.go b/coderd/database/migrations/migrate_test.go index a3ec5435267..1c13210637f 100644 --- a/coderd/database/migrations/migrate_test.go +++ b/coderd/database/migrations/migrate_test.go @@ -5,7 +5,6 @@ import ( "database/sql" "encoding/json" "fmt" - "io/fs" "os" "path/filepath" "slices" @@ -2958,38 +2957,22 @@ func mustJSON(t *testing.T, v any) []byte { func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { t.Parallel() - const migrationVersion = 566 + const previousMigrationVersion = 565 sqlDB := testSQLDB(t) - - // stepperUpToLatest runs the stepper to completion: a Stepper cannot - // stop early, it closes only when the steps are exhausted, and each - // call commits the driver's transaction. The assertion only requires - // that target was APPLIED, not that it is the latest: a stacked PR may - // add a later migration, so encoding "target is latest" would redden - // any such child on its merge ref. - stepperUpToLatest := func(target uint) { - t.Helper() - next, err := migrations.Stepper(sqlDB) + next, err := migrations.Stepper(sqlDB) + require.NoError(t, err) + for { + version, more, err := next() require.NoError(t, err) - last := uint(0) - for { - version, more, err := next() - require.NoError(t, err) - if !more { - break - } - last = version + if !more { + t.Fatalf("migration %d not found", previousMigrationVersion) + } + if version == previousMigrationVersion { + break } - require.GreaterOrEqual(t, last, target, "migration %d must be applied", target) } - // Migrate fully up first. The Stepper cannot stop at an intermediate - // version (it runs to completion and each run is one transaction), so - // fixtures are seeded post-schema and the explosion is exercised as - // DOWN (restore pre-explosion state) then UP (re-explode from it). - stepperUpToLatest(migrationVersion) - ctx := testutil.Context(t, testutil.WaitSuperLong) now := time.Now().UTC().Truncate(time.Microsecond) @@ -3004,9 +2987,7 @@ func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { c3ID := uuid.New() // soft-deleted, referenced in orgB only c4ID := uuid.New() // soft-deleted, unreferenced: never copied c5ID := uuid.New() // live plain config - // aclLateConfigID simulates a config written by an older binary during a - // rolling upgrade, without the everyone ACL entry. - aclLateConfigID := uuid.New() + emptyACLConfigID := uuid.New() chatBID := uuid.New() // chat in orgB pinned to live c1 chatB3ID := uuid.New() // chat in orgB pinned to deleted c3 chatDID := uuid.New() // chat in soft-deleted orgD pinned to live c1 @@ -3070,7 +3051,7 @@ func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { insertConfig(c3ID, "gpt-4-legacy", false, true, everyoneACL) insertConfig(c4ID, "gpt-4-ancient", false, true, everyoneACL) insertConfig(c5ID, "gpt-5.2-nano", false, false, everyoneACL) - insertConfig(aclLateConfigID, "gpt-5.2-late", false, false, `{}`) + insertConfig(emptyACLConfigID, "gpt-5.2-empty-acl", false, false, `{}`) // Chats: orgB pinned to live c1, orgB second chat pinned to deleted c3, // soft-deleted orgD pinned to live c1. @@ -3168,15 +3149,10 @@ func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { ) } - // The fixtures were seeded post-schema in pre-explosion shape - // (default-org configs, references pointing at originals). The Stepper - // cannot stop mid-way, so the explosion of the seeded data is driven by - // rewinding the version row to the predecessor and stepping to latest, - // which re-applies this migration's up over the seeded rows. - _, err := sqlDB.ExecContext(ctx, fmt.Sprintf( - "TRUNCATE schema_migrations; INSERT INTO schema_migrations (version, dirty) VALUES (%d, false)", migrationVersion-1)) + upSQL, err := os.ReadFile("000566_chat_model_config_org_explosion.up.sql") + require.NoError(t, err) + _, err = sqlDB.ExecContext(ctx, string(upSQL)) require.NoError(t, err) - stepperUpToLatest(migrationVersion) // copyID resolves the copy of orig in org by natural attributes (the // migration persists no mapping; this is also what the down relies on). @@ -3319,7 +3295,6 @@ func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { WHERE NOT co.is_default AND NOT co.deleted`).Scan(&dangling)) require.Zero(t, dangling, "no live-org chat references a default-org config") - // --- ACL re-key and backfill --- // Each copy carries only its target organization's everyone entry. for _, orgID := range []uuid.UUID{orgBID, orgCID} { rows, err := sqlDB.QueryContext(ctx, @@ -3335,14 +3310,6 @@ func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { require.NoError(t, rows.Err()) require.NoError(t, rows.Close()) } - // The pre-existing row with an empty group_acl was backfilled. - var aclRaw []byte - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT group_acl FROM chat_model_configs WHERE id = $1", aclLateConfigID).Scan(&aclRaw)) - require.JSONEq(t, string(mustJSON(t, map[string]any{ - defaultOrgID.String(): map[string]any{"permissions": []string{"read"}}, - })), string(aclRaw), "row missing the everyone entry is backfilled") - // --- Threshold fan-out --- // c1 keys fanned out to the orgB and orgC copies for both users (follows // copies, not membership); the c3 key fanned out to the orgB copy only; @@ -3412,12 +3379,7 @@ func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { user1ID).Scan(&overrideValue)) require.Equal(t, "chat_default", overrideValue, "non-threshold keys are untouched") - // --- Down round-trip on the exploded state. The framework's Stepper - // cannot commit after exactly one down step (its driver commits only - // when the stepper exhausts), so the down file is executed directly: - // migrations are plain SQL and this is the same text golang-migrate - // would run. - downSQL, err := fs.ReadFile(migrations.MigrationFS(), "000566_chat_model_config_org_explosion.down.sql") + downSQL, err := os.ReadFile("000566_chat_model_config_org_explosion.down.sql") require.NoError(t, err) _, err = sqlDB.ExecContext(ctx, string(downSQL)) require.NoError(t, err) @@ -3450,22 +3412,15 @@ func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) require.Equal(t, c3ID, msgCfg, "down: orgB deleted-model debug run restored") - // Fanned-out threshold keys are gone; exactly the seeded valid key set - // remains. The two hostile keys (malformed/empty suffix) are pruned as - // dangling: they cannot name an existing config, and the down must not - // abort on their uuid cast. + // The down removes only fanned-out threshold keys. It preserves all + // seeded keys, including malformed and dangling keys. var thresholdCount int require.NoError(t, sqlDB.QueryRowContext(ctx, "SELECT COUNT(*) FROM user_configs WHERE key LIKE 'chat_compaction_threshold_pct:%'").Scan(&thresholdCount)) - require.Equal(t, 4, thresholdCount, "down: only the 4 seeded valid threshold keys remain") + require.Equal(t, 6, thresholdCount, "down: all 6 seeded threshold keys remain") - // Step forward again: the down executed outside the driver, so rewind - // the version row and re-apply the up. Shape re-asserted at count - // level. - _, err = sqlDB.ExecContext(ctx, fmt.Sprintf( - "TRUNCATE schema_migrations; INSERT INTO schema_migrations (version, dirty) VALUES (%d, false)", migrationVersion-1)) + _, err = sqlDB.ExecContext(ctx, string(upSQL)) require.NoError(t, err) - stepperUpToLatest(migrationVersion) assertCount(defaultOrgID, 6, "re-up: default org keeps its originals") assertCount(orgBID, 5, "re-up: orgB copies recreated") assertCount(orgCID, 4, "re-up: orgC copies recreated") diff --git a/coderd/exp_chats.go b/coderd/exp_chats.go index ec08fdc96a7..794203f692d 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -1129,6 +1129,7 @@ func (api *API) getUserChatProviderAvailability( func (api *API) userCanUseChatModelConfig( ctx context.Context, userID uuid.UUID, + organizationID uuid.UUID, modelConfigID uuid.UUID, ) (database.ChatModelConfig, chatModelConfigUnavailableReason, error) { if modelConfigID == uuid.Nil { @@ -1145,7 +1146,7 @@ func (api *API) userCanUseChatModelConfig( } return database.ChatModelConfig{}, chatModelConfigAvailable, err } - if !model.Enabled { + if model.OrganizationID != organizationID || !model.Enabled { return database.ChatModelConfig{}, chatModelConfigUnavailableModelNotFoundOrDisabled, nil } @@ -1176,9 +1177,10 @@ func (api *API) userCanUseChatModelConfig( func (api *API) validateUserChatModelConfigAvailable( ctx context.Context, userID uuid.UUID, + organizationID uuid.UUID, modelConfigID uuid.UUID, ) (database.ChatModelConfig, int, *codersdk.Response) { - modelConfig, reason, err := api.userCanUseChatModelConfig(ctx, userID, modelConfigID) + modelConfig, reason, err := api.userCanUseChatModelConfig(ctx, userID, organizationID, modelConfigID) if err != nil { return database.ChatModelConfig{}, http.StatusInternalServerError, &codersdk.Response{ Message: "Internal error validating model config override.", @@ -1219,12 +1221,13 @@ func (api *API) validateUserChatModelConfigAvailable( func (api *API) validateExplicitChatModelConfigAvailable( ctx context.Context, userID uuid.UUID, + organizationID uuid.UUID, modelConfigID uuid.UUID, ) (int, *codersdk.Response) { if modelConfigID == uuid.Nil { return 0, nil } - _, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, modelConfigID) + _, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, organizationID, modelConfigID) return status, resp } @@ -2803,7 +2806,7 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) { if req.ModelConfigID != nil { modelConfigID = *req.ModelConfigID } - if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, modelConfigID); resp != nil { + if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, chat.OrganizationID, modelConfigID); resp != nil { httpapi.Write(ctx, rw, status, *resp) return } @@ -2995,7 +2998,7 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) { if req.ModelConfigID != nil { editModelConfigID = *req.ModelConfigID } - if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, editModelConfigID); resp != nil { + if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, chat.OrganizationID, editModelConfigID); resp != nil { httpapi.Write(ctx, rw, status, *resp) return } @@ -4396,7 +4399,7 @@ func (api *API) resolveCreateChatModelConfigID( Message: "Invalid model config ID.", } } - if _, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, *req.ModelConfigID); resp != nil { + if _, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, req.OrganizationID, *req.ModelConfigID); resp != nil { return uuid.Nil, nil, status, resp } return *req.ModelConfigID, nil, 0, nil @@ -4410,7 +4413,7 @@ func (api *API) resolveCreateChatModelConfigID( } } if !personalOverridesEnabled { - id, status, resp := api.defaultCreateChatModelConfigID(ctx) + id, status, resp := api.defaultCreateChatModelConfigID(ctx, req.OrganizationID) return id, nil, status, resp } @@ -4445,6 +4448,7 @@ func (api *API) resolveCreateChatModelConfigID( _, reason, err := api.userCanUseChatModelConfig( ctx, userID, + req.OrganizationID, parsed.ModelConfigID, ) if err != nil { @@ -4473,26 +4477,15 @@ func (api *API) resolveCreateChatModelConfigID( } } - id, status, resp := api.defaultCreateChatModelConfigID(ctx) + id, status, resp := api.defaultCreateChatModelConfigID(ctx, req.OrganizationID) return id, nil, status, resp } func (api *API) defaultCreateChatModelConfigID( ctx context.Context, + organizationID uuid.UUID, ) (uuid.UUID, int, *codersdk.Response) { - // The request carries a validated organization, but the pre-cutover - // handler deliberately resolves the deployment-default model until the - // API layer completes the organization-scoping cutover. The default-org - // lookup is internal and does not expose organization data to the user. - //nolint:gocritic // Internal default-org resolution, scoped to this call. - defaultOrg, err := api.Database.GetDefaultOrganization(dbauthz.AsChatd(ctx)) - if err != nil { - return uuid.Nil, http.StatusInternalServerError, &codersdk.Response{ - Message: "Failed to resolve chat model config.", - Detail: err.Error(), - } - } - defaultModelConfig, err := api.Database.GetDefaultChatModelConfig(ctx, defaultOrg.ID) + defaultModelConfig, err := api.Database.GetDefaultChatModelConfig(ctx, organizationID) if err != nil { if xerrors.Is(err, sql.ErrNoRows) { return uuid.Nil, http.StatusBadRequest, &codersdk.Response{ @@ -5029,7 +5022,18 @@ func (api *API) putUserChatPersonalModelOverride(rw http.ResponseWriter, r *http }) return } - modelConfig, status, resp := api.validateUserChatModelConfigAvailable(ctx, apiKey.UserID, parsedModelConfigID) + // Personal model overrides are user-global. They select from the + // default organization's model configs. + //nolint:gocritic // This lookup resolves internal deployment configuration. + defaultOrg, err := api.Database.GetDefaultOrganization(dbauthz.AsChatd(ctx)) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Internal error validating model config override.", + Detail: err.Error(), + }) + return + } + modelConfig, status, resp := api.validateUserChatModelConfigAvailable(ctx, apiKey.UserID, defaultOrg.ID, parsedModelConfigID) if resp != nil { httpapi.Write(ctx, rw, status, *resp) return diff --git a/coderd/exp_chats_test.go b/coderd/exp_chats_test.go index d23811a3bb3..95bc72c9450 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -327,32 +327,83 @@ func insertAssistantMessage( func TestPostChats(t *testing.T) { t.Parallel() - t.Run("SuccessNonDefaultOrgUsesDeploymentDefault", func(t *testing.T) { + t.Run("SuccessNonDefaultOrgUsesOrgDefault", func(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitLong) client, db := newChatClientWithDatabase(t) _ = coderdtest.CreateFirstUser(t, client.Client) - modelConfig := createChatModelConfig(t, client) - - // A member of a non-default org cannot read the default - // organization object, but omitting model_config_id must still - // resolve the deployment default while every config lives there. + defaultConfig := createChatModelConfig(t, client) org := dbgen.Organization(t, db, database.Organization{IsDefault: false}) + orgConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: defaultConfig.AIProviderID, Valid: true}, + Model: "org-default-" + uuid.NewString(), + Enabled: true, + IsDefault: true, + OrganizationID: org.ID, + }) memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, org.ID, rbac.ScopedRoleAgentsAccess(org.ID)) memberClient := codersdk.NewExperimentalClient(memberClientRaw) chat, err := memberClient.CreateChat(ctx, codersdk.CreateChatRequest{ OrganizationID: org.ID, - Content: []codersdk.ChatInputPart{ - { - Type: codersdk.ChatInputPartTypeText, - Text: "hello from a non-default org", - }, - }, + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "hello from a non-default org", + }}, }) require.NoError(t, err) - require.Equal(t, modelConfig.ID, chat.LastModelConfigID) + require.Equal(t, orgConfig.ID, chat.LastModelConfigID) + }) + + t.Run("NonDefaultOrgWithoutDefaultRejected", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + _ = coderdtest.CreateFirstUser(t, client.Client) + _ = createChatModelConfig(t, client) + org := dbgen.Organization(t, db, database.Organization{IsDefault: false}) + memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, org.ID, rbac.ScopedRoleAgentsAccess(org.ID)) + memberClient := codersdk.NewExperimentalClient(memberClientRaw) + + _, err := memberClient.CreateChat(ctx, codersdk.CreateChatRequest{ + OrganizationID: org.ID, + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "no model is configured", + }}, + }) + sdkErr := requireSDKError(t, err, http.StatusBadRequest) + require.Equal(t, "No default chat model config is configured.", sdkErr.Message) + }) + + t.Run("CrossOrgExplicitModelRejected", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + defaultConfig := createChatModelConfig(t, client) + otherOrg := dbgen.Organization(t, db, database.Organization{IsDefault: false}) + otherConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: defaultConfig.AIProviderID, Valid: true}, + Model: "cross-org-" + uuid.NewString(), + Enabled: true, + IsDefault: true, + OrganizationID: otherOrg.ID, + }) + + _, err := client.CreateChat(ctx, codersdk.CreateChatRequest{ + OrganizationID: firstUser.OrganizationID, + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "reject another organization's model", + }}, + ModelConfigID: ptr.Ref(otherConfig.ID), + }) + sdkErr := requireSDKError(t, err, http.StatusBadRequest) + require.Equal(t, "Invalid model_config_id: model config not found or disabled.", sdkErr.Message) }) t.Run("Success", func(t *testing.T) { @@ -7159,6 +7210,40 @@ func TestPostChatMessages(t *testing.T) { require.Equal(t, "Invalid model_config_id: provider is not enabled for this model.", sdkErr.Message) }) + t.Run("CrossOrgModelConfigRejected", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + defaultConfig := createChatModelConfig(t, client) + chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{ + OrganizationID: firstUser.OrganizationID, + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "initial message before cross-org switch", + }}, + }) + require.NoError(t, err) + + otherOrg := dbgen.Organization(t, db, database.Organization{IsDefault: false}) + otherConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: defaultConfig.AIProviderID, Valid: true}, + Model: "cross-org-send-" + uuid.NewString(), + Enabled: true, + OrganizationID: otherOrg.ID, + }) + _, err = client.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{ + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "reject another organization's model", + }}, + ModelConfigID: ptr.Ref(otherConfig.ID), + }) + sdkErr := requireSDKError(t, err, http.StatusBadRequest) + require.Equal(t, "Invalid model_config_id: model config not found or disabled.", sdkErr.Message) + }) + t.Run("ProviderDisabledDefaultFallbackRejected", func(t *testing.T) { t.Parallel() @@ -8674,6 +8759,47 @@ func TestPatchChatMessage(t *testing.T) { require.False(t, foundOriginalInChat) }) + t.Run("CrossOrgModelConfigRejected", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + client, db := newChatClientWithDatabase(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + defaultConfig := createChatModelConfig(t, client) + chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{ + OrganizationID: firstUser.OrganizationID, + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "before cross-org edit", + }}, + }) + require.NoError(t, err) + messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil) + require.NoError(t, err) + userMessageID := messagesResult.Messages[0].ID + + otherOrg := dbgen.Organization(t, db, database.Organization{IsDefault: false}) + otherConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: defaultConfig.AIProviderID, Valid: true}, + Model: "cross-org-edit-" + uuid.NewString(), + Enabled: true, + OrganizationID: otherOrg.ID, + }) + _, err = client.EditChatMessage(ctx, chat.ID, userMessageID, codersdk.EditChatMessageRequest{ + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "reject another organization's model", + }}, + ModelConfigID: ptr.Ref(otherConfig.ID), + }) + sdkErr := requireSDKError(t, err, http.StatusBadRequest) + require.Equal(t, "Invalid model_config_id: model config not found or disabled.", sdkErr.Message) + + storedChat, err := db.GetChatByID(dbauthz.AsSystemRestricted(ctx), chat.ID) + require.NoError(t, err) + require.Equal(t, defaultConfig.ID, storedChat.LastModelConfigID) + }) + t.Run("ReasoningEffort", func(t *testing.T) { t.Parallel() @@ -13801,6 +13927,35 @@ func TestCreateChatPersonalModelOverrideRoot(t *testing.T) { require.Equal(t, ptr.Ref("high"), chat.LastReasoningEffort) }) + t.Run("CrossOrgRootModelFallsBackToOrgDefault", func(t *testing.T) { + org := dbgen.Organization(t, db, database.Organization{IsDefault: false}) + orgModel := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: defaultModel.AIProviderID, Valid: true}, + Model: "org-root-personal-" + uuid.NewString(), + Enabled: true, + IsDefault: true, + OrganizationID: org.ID, + }) + otherClientRaw, otherUser := coderdtest.CreateAnotherUser( + t, + adminClient.Client, + org.ID, + rbac.ScopedRoleAgentsAccess(org.ID), + ) + otherClient := codersdk.NewExperimentalClient(otherClientRaw) + upsertRootRaw(otherUser.ID, "model:"+overrideModel.ID.String()) + + chat, err := otherClient.CreateChat(ctx, codersdk.CreateChatRequest{ + OrganizationID: org.ID, + Content: []codersdk.ChatInputPart{{ + Type: codersdk.ChatInputPartTypeText, + Text: "cross-org root model falls back", + }}, + }) + require.NoError(t, err) + require.Equal(t, orgModel.ID, chat.LastModelConfigID) + }) + t.Run("UnavailableRootModelFallsBackToDefault", func(t *testing.T) { upsertRootRaw(firstUser.UserID, "model:"+disabledModel.ID.String()) chat := createChat(adminClient, "disabled root model falls back", nil) diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 8334fd67c13..4bf3c39dfb4 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -1237,7 +1237,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C deploymentPrompt := p.resolveDeploymentSystemPrompt(ctx) if opts.ModelConfigID != uuid.Nil { - if err := requireEnabledChatModelConfig(ctx, p.db, opts.ModelConfigID); err != nil { + if err := requireEnabledChatModelConfig(ctx, p.db, opts.OrganizationID, opts.ModelConfigID); err != nil { return database.Chat{}, err } } @@ -1251,7 +1251,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C contentParts := opts.InitialUserContent if p.hooks.Enabled() { // Validate model admission before dispatch, matching the insert path. - if err := validateCreateModelConfigID(ctx, p.db, opts.ModelConfigID); err != nil { + if err := validateCreateModelConfigID(ctx, p.db, opts.OrganizationID, opts.ModelConfigID); err != nil { return database.Chat{}, err } turnID := uuid.New() @@ -1549,7 +1549,7 @@ func resolveSendMessageModelConfigID( return resolveFallbackModelConfigID(ctx, store, chat.OrganizationID, chat.LastModelConfigID) } - if err := requireEnabledChatModelConfig(ctx, store, requested); err != nil { + if err := requireEnabledChatModelConfig(ctx, store, chat.OrganizationID, requested); err != nil { return uuid.Nil, err } return requested, nil @@ -1560,10 +1560,12 @@ func resolveSendMessageModelConfigID( func requireEnabledChatModelConfig( ctx context.Context, store database.Store, + organizationID uuid.UUID, modelConfigID uuid.UUID, ) error { chatdCtx := chatdModelConfigLookupContext(ctx) - if _, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID); err != nil { + modelConfig, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID) + if err != nil { if errors.Is(err, sql.ErrNoRows) { return xerrors.Errorf( "%w: %s", @@ -1577,20 +1579,27 @@ func requireEnabledChatModelConfig( err, ) } + if modelConfig.OrganizationID != organizationID { + return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID) + } return nil } -func validateCreateModelConfigID(ctx context.Context, store database.Store, modelConfigID uuid.UUID) error { +func validateCreateModelConfigID(ctx context.Context, store database.Store, organizationID, modelConfigID uuid.UUID) error { if modelConfigID == uuid.Nil { return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID) } chatdCtx := chatdModelConfigLookupContext(ctx) - if _, err := store.GetChatModelConfigByID(chatdCtx, modelConfigID); err != nil { + modelConfig, err := store.GetChatModelConfigByID(chatdCtx, modelConfigID) + if err != nil { if errors.Is(err, sql.ErrNoRows) { return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID) } return xerrors.Errorf("get requested model config %s: %w", modelConfigID, err) } + if modelConfig.OrganizationID != organizationID { + return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID) + } return nil } @@ -1602,8 +1611,10 @@ func resolveFallbackModelConfigID( ) (uuid.UUID, error) { chatdCtx := chatdModelConfigLookupContext(ctx) if modelConfigID != uuid.Nil { - if _, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID); err == nil { - return modelConfigID, nil + if modelConfig, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID); err == nil { + if modelConfig.OrganizationID == organizationID { + return modelConfigID, nil + } } else if !errors.Is(err, sql.ErrNoRows) { return uuid.Nil, xerrors.Errorf( "get chat model config %s: %w", @@ -1613,7 +1624,7 @@ func resolveFallbackModelConfigID( } } - defaultConfig, err := defaultChatModelConfigForOrg(chatdCtx, store, organizationID) + defaultConfig, err := store.GetDefaultChatModelConfig(chatdCtx, organizationID) if err != nil { if errors.Is(err, sql.ErrNoRows) { return uuid.Nil, ErrNoDefaultChatModelConfig @@ -1641,12 +1652,13 @@ func resolveFallbackModelConfigID( func validateModelConfigOverride( ctx context.Context, store database.Store, + organizationID uuid.UUID, requested uuid.UUID, ) (uuid.NullUUID, error) { if requested == uuid.Nil { return uuid.NullUUID{}, nil } - if err := requireEnabledChatModelConfig(ctx, store, requested); err != nil { + if err := requireEnabledChatModelConfig(ctx, store, organizationID, requested); err != nil { return uuid.NullUUID{}, err } return uuid.NullUUID{UUID: requested, Valid: true}, nil @@ -1669,87 +1681,6 @@ func validateEditTarget(ctx context.Context, store database.Store, chatID uuid.U return nil } -// defaultChatModelConfigForOrg resolves the default model config that -// serves organizationID. An org that owns no configs of its own reads -// the default org's default instead, which preserves the -// pre-org-scoping behavior where every chat saw the deployment-wide -// configs. The default org never falls back: a miss there is a real -// absence and returns sql.ErrNoRows. -// -// The returned config's OrganizationID identifies which org's configs -// apply, so callers that need the whole list can read it from there. -// TODO(mafredri): remove after CODAGT-709 M3 (org-scoping cutover); -// per-org resolution becomes strict. -func defaultChatModelConfigForOrg( - ctx context.Context, - store database.Store, - organizationID uuid.UUID, -) (database.ChatModelConfig, error) { - config, err := store.GetDefaultChatModelConfig(ctx, organizationID) - if err == nil { - return config, nil - } - if !errors.Is(err, sql.ErrNoRows) { - return database.ChatModelConfig{}, err - } - defaultOrg, err := store.GetDefaultOrganization(ctx) - if err != nil { - return database.ChatModelConfig{}, xerrors.Errorf("get default organization: %w", err) - } - if defaultOrg.ID == organizationID { - return database.ChatModelConfig{}, sql.ErrNoRows - } - return store.GetDefaultChatModelConfig(ctx, defaultOrg.ID) -} - -// enabledChatModelConfigsWithDefaultOrgFallback returns the organization's -// enabled configs. It uses the default org's configs when the organization -// owns no configs, which preserves the pre-org-scoping behavior. The default -// org itself never falls back. -// -// An empty initial result does not prove the organization owns no configs. Its -// configs can all be disabled or use disabled providers. The resolved default -// config is the ownership marker because every write path preserves one default -// in each organization that owns configs. -// TODO(mafredri): remove after CODAGT-709 M3 (org-scoping cutover); -// organizations list strictly within their own configs. -func enabledChatModelConfigsWithDefaultOrgFallback( - ctx context.Context, - store database.Store, - organizationID uuid.UUID, -) ([]database.GetEnabledChatModelConfigsByOrganizationRow, error) { - rows, err := store.GetEnabledChatModelConfigsByOrganization(ctx, organizationID) - if err != nil { - return nil, err - } - if len(rows) > 0 { - return rows, nil - } - - defaultConfig, err := defaultChatModelConfigForOrg(ctx, store, organizationID) - if err != nil { - if errors.Is(err, sql.ErrNoRows) { - return rows, nil - } - return nil, xerrors.Errorf("resolve default chat model config: %w", err) - } - if defaultConfig.OrganizationID == organizationID { - // The default can appear after the initial list read. Re-fetch the - // organization's enabled configs to avoid returning a stale empty list. - rows, err = store.GetEnabledChatModelConfigsByOrganization(ctx, organizationID) - if err != nil { - return nil, xerrors.Errorf("re-fetch organization enabled chat model configs: %w", err) - } - return rows, nil - } - - fallbackRows, err := store.GetEnabledChatModelConfigsByOrganization(ctx, defaultConfig.OrganizationID) - if err != nil { - return nil, xerrors.Errorf("get default org enabled chat model configs: %w", err) - } - return fallbackRows, nil -} - // EditMessage replaces an earlier user message and discards the // active-history suffix through chatstate.EditMessage. Model-config // override validation and usage-limit admission run in the same @@ -1783,7 +1714,7 @@ func (p *Server) EditMessage( if err := validateEditTarget(ctx, p.db, opts.ChatID, opts.EditedMessageID); err != nil { return EditMessageResult{}, err } - if _, err := validateModelConfigOverride(ctx, p.db, opts.ModelConfigID); err != nil { + if _, err := validateModelConfigOverride(ctx, p.db, chat.OrganizationID, opts.ModelConfigID); err != nil { return EditMessageResult{}, err } sessionStartHookResult, err = p.hooks.Trigger(ctx, chathooks.ChatFor(chat, &turnID), chathooks.Message{Source: chathooks.SessionStartSourceClear}, agenthooks.EventSessionStart, dispatch.CapacityClassAdmission) @@ -1843,7 +1774,7 @@ func (p *Server) EditMessage( } editedMsg = target - modelOverride, err := validateModelConfigOverride(ctx, store, opts.ModelConfigID) + modelOverride, err := validateModelConfigOverride(ctx, store, lockedChat.OrganizationID, opts.ModelConfigID) if err != nil { return err } @@ -2829,7 +2760,7 @@ func (p *Server) resolveManualTitleModel( return overrideModel, overrideConfig, nil } - configs, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, store, chat.OrganizationID) + configs, err := store.GetEnabledChatModelConfigsByOrganization(ctx, chat.OrganizationID) if err != nil { p.logger.Debug(ctx, "failed to list manual title model configs", slog.F("chat_id", chat.ID), diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index 24bd6b2e45d..1d5ae1bf3bc 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -1129,13 +1129,9 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) { }, ).Return(nil, nil) db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil) + // An empty org list falls through to the chat's fallback model; + // strict org scoping reads no other org's configs. db.EXPECT().GetEnabledChatModelConfigsByOrganization(gomock.Any(), gomock.Any()).Return(nil, nil) - // An empty org list only triggers the pre-cutover fallback when the - // org owns no configs at all, which the missing default proves. The - // default org has no default either, so the empty list stands. - db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), gomock.Any()). - Return(database.ChatModelConfig{}, sql.ErrNoRows).Times(2) - db.EXPECT().GetDefaultOrganization(gomock.Any()).Return(database.Organization{ID: uuid.New()}, nil) db.EXPECT().InTx(gomock.Any(), nil).DoAndReturn( func(fn func(database.Store) error, opts *database.TxOptions) error { @@ -1280,13 +1276,9 @@ func TestRegenerateChatTitle_SkipsPersistWhenTitleChangedConcurrently(t *testing }, ).Return(nil, nil) db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil) + // An empty org list falls through to the chat's fallback model; + // strict org scoping reads no other org's configs. db.EXPECT().GetEnabledChatModelConfigsByOrganization(gomock.Any(), gomock.Any()).Return(nil, nil) - // An empty org list only triggers the pre-cutover fallback when the - // org owns no configs at all, which the missing default proves. The - // default org has no default either, so the empty list stands. - db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), gomock.Any()). - Return(database.ChatModelConfig{}, sql.ErrNoRows).Times(2) - db.EXPECT().GetDefaultOrganization(gomock.Any()).Return(database.Organization{ID: uuid.New()}, nil) db.EXPECT().InTx(gomock.Any(), nil).DoAndReturn( func(fn func(database.Store) error, _ *database.TxOptions) error { @@ -3805,41 +3797,38 @@ func TestResolveFallbackModelConfigID(t *testing.T) { require.Equal(t, defaultModel.ID, resolved) }) - t.Run("NonDefaultOrgFallsBackToDefaultOrgDefault", func(t *testing.T) { + t.Run("NonDefaultOrgWithoutOwnDefaultMisses", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) ctx := testutil.Context(t, testutil.WaitShort) - // The chat's org has no configs of its own; the deployment - // default lives in the default org. Pre-cutover behavior must - // resolve it for chats in any org. + // The chat's org has no configs of its own; the default org has + // one. Strict scoping resolves configs only within the chat's + // org, so the lookup reports no default. otherOrgID := newModelConfigOrg(t, db) defaultOrg, err := db.GetDefaultOrganization(ctx) require.NoError(t, err) provider := newProvider(t, db, true) - defaultModel := newModelConfig(t, db, defaultOrg.ID, provider.ID, true) + _ = newModelConfig(t, db, defaultOrg.ID, provider.ID, true) - resolved, err := resolveFallbackModelConfigID(ctx, db, otherOrgID, uuid.Nil) - require.NoError(t, err) - require.Equal(t, defaultModel.ID, resolved) + _, err = resolveFallbackModelConfigID(ctx, db, otherOrgID, uuid.Nil) + require.ErrorIs(t, err, ErrNoDefaultChatModelConfig) }) - t.Run("FallbackReadsDefaultOrgAsChatd", func(t *testing.T) { + t.Run("OrgDefaultResolvesAsChatd", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) ctx := testutil.Context(t, testutil.WaitShort) - // The pre-cutover fallback reads the default organization under - // the chatd subject, which must be authorized to read - // organizations or every fallback path fails closed. - otherOrgID := newModelConfigOrg(t, db) - defaultOrg, err := db.GetDefaultOrganization(ctx) - require.NoError(t, err) + // The fallback reads the org default under the chatd subject, + // which must be authorized to read chat model configs or every + // fallback path fails closed. + orgID := newModelConfigOrg(t, db) provider := newProvider(t, db, true) - defaultModel := newModelConfig(t, db, defaultOrg.ID, provider.ID, true) + defaultModel := newModelConfig(t, db, orgID, provider.ID, true) chatdCtx := dbauthz.AsChatd(ctx) - resolved, err := resolveFallbackModelConfigID(chatdCtx, db, otherOrgID, uuid.Nil) + resolved, err := resolveFallbackModelConfigID(chatdCtx, db, orgID, uuid.Nil) require.NoError(t, err) require.Equal(t, defaultModel.ID, resolved) }) @@ -3867,11 +3856,25 @@ func TestResolveFallbackModelConfigID(t *testing.T) { provider := newProvider(t, db, true) model := newModelConfig(t, db, orgID, provider.ID, false) - resolved, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{}, model.ID) + resolved, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{OrganizationID: orgID}, model.ID) require.NoError(t, err) require.Equal(t, model.ID, resolved) }) + t.Run("ExplicitCrossOrgModelRejected", func(t *testing.T) { + t.Parallel() + db, _ := dbtestutil.NewDB(t) + ctx := testutil.Context(t, testutil.WaitShort) + + chatOrgID := newModelConfigOrg(t, db) + modelOrgID := newModelConfigOrg(t, db) + provider := newProvider(t, db, true) + model := newModelConfig(t, db, modelOrgID, provider.ID, false) + + _, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{OrganizationID: chatOrgID}, model.ID) + require.ErrorIs(t, err, ErrInvalidModelConfigID) + }) + // An explicit model whose provider was disabled after the coderd // preflight must still be rejected inside the daemon. t.Run("ExplicitProviderDisabledRejected", func(t *testing.T) { @@ -3883,7 +3886,7 @@ func TestResolveFallbackModelConfigID(t *testing.T) { disabledProvider := newProvider(t, db, false) model := newModelConfig(t, db, orgID, disabledProvider.ID, false) - _, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{}, model.ID) + _, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{OrganizationID: orgID}, model.ID) require.ErrorIs(t, err, ErrInvalidModelConfigID) }) diff --git a/coderd/x/chatd/compaction_override.go b/coderd/x/chatd/compaction_override.go index fc764ab58db..8068da64613 100644 --- a/coderd/x/chatd/compaction_override.go +++ b/coderd/x/chatd/compaction_override.go @@ -89,6 +89,9 @@ func (p *Server) resolveCompactionOverrideConfig( if err != nil || !overrideSet { return nil, err } + if modelConfig.OrganizationID != chat.OrganizationID { + return nil, err + } // Already validated by the shared resolver; failure is unreachable. resolvedProvider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint( modelConfig.Model, diff --git a/coderd/x/chatd/compaction_override_internal_test.go b/coderd/x/chatd/compaction_override_internal_test.go index 166263c3d52..24a94f4b0f9 100644 --- a/coderd/x/chatd/compaction_override_internal_test.go +++ b/coderd/x/chatd/compaction_override_internal_test.go @@ -146,6 +146,34 @@ func TestResolveCompactionOverrideConfig_DisabledConfigFallsBack(t *testing.T) { require.Nil(t, override) } +func TestResolveCompactionOverrideConfig_CrossOrgFallsBack(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + chat, _ := titleOverrideTestChatAndMessages(t) + chat.OrganizationID = uuid.New() + overrideConfig := titleOverrideModelConfig("gpt-4.1", true) + overrideConfig.OrganizationID = uuid.New() + providerID := uuid.New() + overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + + db.EXPECT().GetChatCompactionModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) + db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "test-key", + }}, nil).AnyTimes() + + server := titleOverrideTestServer(db, logger) + override, err := server.resolveCompactionOverrideConfig(ctx, chat) + require.NoError(t, err) + require.Nil(t, override) +} + func TestResolveCompactionOverrideConfig_MissingCredentialsFallsBack(t *testing.T) { t.Parallel() diff --git a/coderd/x/chatd/configcache.go b/coderd/x/chatd/configcache.go index 213905e1be6..a1e9cb34fef 100644 --- a/coderd/x/chatd/configcache.go +++ b/coderd/x/chatd/configcache.go @@ -286,11 +286,8 @@ func (c *chatConfigCache) storeModelConfig(snap modelConfigSnapshot, config data } // DefaultModelConfig returns the default model config for the given -// organization. Until the org-scoping cutover, an org without its own -// default falls back to the default org's config. -// TODO(mafredri): remove after CODAGT-709 M3 (org-scoping cutover); -// orgs resolve strictly within their own configs after the org-scoping -// cutover. +// organization. Orgs resolve strictly within their own configs; an org +// without its own default reports absence. func (c *chatConfigCache) DefaultModelConfig(ctx context.Context, orgID uuid.UUID) (database.ChatModelConfig, error) { if config, ok := c.cachedDefaultModelConfig(orgID); ok { return config, nil @@ -302,7 +299,7 @@ func (c *chatConfigCache) DefaultModelConfig(ctx context.Context, orgID uuid.UUI return cached, nil } - fetched, err := defaultChatModelConfigForOrg(c.ctx, c.db, orgID) + fetched, err := c.db.GetDefaultChatModelConfig(c.ctx, orgID) if err != nil { return database.ChatModelConfig{}, err } @@ -423,8 +420,8 @@ func (c *chatConfigCache) InvalidateModelConfig(id uuid.UUID) { delete(c.modelConfigs, id) c.modelTopologyEpoch++ // Coarse invalidation: the event does not say whether the changed - // config was a default, nor which orgs resolve to it through the - // default-org fallback, so every per-org default is dropped. + // config was a default, nor for which org, so every per-org default + // is dropped. clear(c.defaultModelConfigs) c.defaultModelConfigGeneration++ c.mu.Unlock() diff --git a/coderd/x/chatd/configcache_internal_test.go b/coderd/x/chatd/configcache_internal_test.go index 92da3c2bbf2..9b53787b4f3 100644 --- a/coderd/x/chatd/configcache_internal_test.go +++ b/coderd/x/chatd/configcache_internal_test.go @@ -336,91 +336,44 @@ func TestConfigCache_DefaultModelConfig_PerOrgKeying(t *testing.T) { require.Equal(t, int32(2), store.defaultModelConfigCallCount(orgB)) } -func TestConfigCache_DefaultModelConfig_DefaultOrgFallback(t *testing.T) { +func TestConfigCache_DefaultModelConfig_CrossOrgIsolation(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitShort) clock := quartz.NewMock(t) - defaultOrgID := uuid.New() otherOrgID := uuid.New() defaultOrgConfig := testChatModelConfig(uuid.New(), "default-org-model") store := &stubChatConfigStore{} store.getDefaultChatModelConfig = func(_ context.Context, orgID uuid.UUID) (database.ChatModelConfig, error) { - if orgID == defaultOrgID { - return defaultOrgConfig, nil - } - return database.ChatModelConfig{}, sql.ErrNoRows - } - store.getDefaultOrganization = func(context.Context) (database.Organization, error) { - return database.Organization{ID: defaultOrgID}, nil - } - cache := newChatConfigCache(ctx, store, clock) - - // An org without its own default resolves the default org's config, - // cached under its own org key. - resolved, err := cache.DefaultModelConfig(ctx, otherOrgID) - require.NoError(t, err) - require.Equal(t, defaultOrgConfig, resolved) - resolvedAgain, err := cache.DefaultModelConfig(ctx, otherOrgID) - require.NoError(t, err) - require.Equal(t, defaultOrgConfig, resolvedAgain) - require.Equal(t, int32(1), store.defaultModelConfigCallCount(otherOrgID)) - require.Equal(t, int32(1), store.defaultModelConfigCallCount(defaultOrgID)) - require.Equal(t, int32(1), store.defaultOrganizationCall.Load()) - - // The default org itself gets no fallback. - _, err = cache.DefaultModelConfig(ctx, defaultOrgID) - require.NoError(t, err) - require.Equal(t, int32(2), store.defaultModelConfigCallCount(defaultOrgID)) - require.Equal(t, int32(1), store.defaultOrganizationCall.Load()) -} - -func TestConfigCache_DefaultModelConfig_DefaultOrgMiss(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - clock := quartz.NewMock(t) - defaultOrgID := uuid.New() - store := &stubChatConfigStore{} - store.getDefaultChatModelConfig = func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { + // The chat's org resolves no default; the default org has one. return database.ChatModelConfig{}, sql.ErrNoRows } - store.getDefaultOrganization = func(context.Context) (database.Organization, error) { - return database.Organization{ID: defaultOrgID}, nil - } cache := newChatConfigCache(ctx, store, clock) - // A miss inside the default org must not recurse into the fallback: - // the default-org resolution happens once, the self-check short - // circuits, and the original miss propagates. - _, err := cache.DefaultModelConfig(ctx, defaultOrgID) + // Orgs resolve strictly within their own configs: an org without its + // own default reports absence and never reads another org's default. + _, err := cache.DefaultModelConfig(ctx, otherOrgID) require.ErrorIs(t, err, sql.ErrNoRows) - require.Equal(t, int32(1), store.defaultOrganizationCall.Load()) - require.Equal(t, int32(1), store.defaultModelConfigCallCount(defaultOrgID)) + require.Equal(t, int32(1), store.defaultModelConfigCallCount(otherOrgID)) + require.Equal(t, int32(0), store.defaultModelConfigCallCount(defaultOrgConfig.OrganizationID)) } -func TestConfigCache_DefaultModelConfig_FallbackMiss(t *testing.T) { +func TestConfigCache_DefaultModelConfig_Miss(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitShort) clock := quartz.NewMock(t) - defaultOrgID := uuid.New() - otherOrgID := uuid.New() + orgID := uuid.New() store := &stubChatConfigStore{} store.getDefaultChatModelConfig = func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { return database.ChatModelConfig{}, sql.ErrNoRows } - store.getDefaultOrganization = func(context.Context) (database.Organization, error) { - return database.Organization{ID: defaultOrgID}, nil - } cache := newChatConfigCache(ctx, store, clock) - // Neither the org nor the default org has a default: the miss - // propagates unchanged. - _, err := cache.DefaultModelConfig(ctx, otherOrgID) + // A miss inside the org propagates unchanged. + _, err := cache.DefaultModelConfig(ctx, orgID) require.ErrorIs(t, err, sql.ErrNoRows) - require.Equal(t, int32(1), store.defaultModelConfigCallCount(otherOrgID)) - require.Equal(t, int32(1), store.defaultModelConfigCallCount(defaultOrgID)) + require.Equal(t, int32(1), store.defaultModelConfigCallCount(orgID)) } func TestConfigCache_UserPrompt_NegativeCaching(t *testing.T) { diff --git a/coderd/x/chatd/subagent.go b/coderd/x/chatd/subagent.go index 04a560a48f7..e115eb56b34 100644 --- a/coderd/x/chatd/subagent.go +++ b/coderd/x/chatd/subagent.go @@ -630,7 +630,7 @@ func (p *Server) listSpawnableModelConfigs( ) ([]map[string]any, error) { //nolint:gocritic // Chatd needs its scoped config and user-data access here. chatdCtx := dbauthz.AsChatd(ctx) - rows, err := enabledChatModelConfigsWithDefaultOrgFallback(chatdCtx, p.db, organizationID) + rows, err := p.db.GetEnabledChatModelConfigsByOrganization(chatdCtx, organizationID) if err != nil { return nil, xerrors.Errorf("get enabled chat model configs: %w", err) } @@ -1265,6 +1265,16 @@ func (p *Server) createChildSubagentChatWithOptions( if opts.modelConfigIDOverride != nil { modelConfigID = *opts.modelConfigIDOverride } + if modelConfigID != uuid.Nil && modelConfigID != parent.LastModelConfigID { + modelConfig, err := p.configCache.ModelConfigByID(ctx, modelConfigID) + if err != nil { + return database.Chat{}, xerrors.Errorf("get child model config: %w", err) + } + if modelConfig.OrganizationID != parent.OrganizationID { + modelConfigID = parent.LastModelConfigID + opts.reasoningEffortOverride = nil + } + } if modelConfigID == uuid.Nil { return database.Chat{}, xerrors.New("model config is required") } diff --git a/coderd/x/chatd/subagent_internal_test.go b/coderd/x/chatd/subagent_internal_test.go index 5ecd505ff94..512f901716c 100644 --- a/coderd/x/chatd/subagent_internal_test.go +++ b/coderd/x/chatd/subagent_internal_test.go @@ -1554,6 +1554,40 @@ func TestCreateChildSubagentChat_OverrideWorksWhenParentHasNoModel(t *testing.T) require.Equal(t, overrideModel.ID, childChat.LastModelConfigID) } +func TestCreateChildSubagentChat_CrossOrgOverrideFallsBackToParent(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{}) + + ctx := chatdTestContext(t) + user, org, parentModel := seedInternalChatDeps(t, db) + otherOrg := dbgen.Organization(t, db, database.Organization{}) + overrideModel := insertInternalChatModelConfig( + t, db, otherOrg.ID, "cross-org-child-"+uuid.NewString(), true, + ) + parentChat := createInternalParentChat( + ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-cross-org-override", + ) + + child, err := server.createChildSubagentChatWithOptions( + ctx, + parentChat, + "delegate work", + "", + childSubagentChatOptions{ + modelConfigIDOverride: &overrideModel.ID, + reasoningEffortOverride: ptr.Ref("high"), + }, + ) + require.NoError(t, err) + + childChat, err := db.GetChatByID(ctx, child.ID) + require.NoError(t, err) + require.Equal(t, parentModel.ID, childChat.LastModelConfigID) + require.False(t, childChat.LastReasoningEffort.Valid) +} + func TestSpawnAgent_ExplicitModelConfigID(t *testing.T) { t.Parallel() @@ -4859,205 +4893,7 @@ func TestListAgents(t *testing.T) { }) } -type enabledChatModelConfigsReadSkewStore struct { - database.Store - organizationID uuid.UUID - config database.ChatModelConfig - listCalls atomic.Int32 -} - -func (s *enabledChatModelConfigsReadSkewStore) GetEnabledChatModelConfigsByOrganization( - _ context.Context, - organizationID uuid.UUID, -) ([]database.GetEnabledChatModelConfigsByOrganizationRow, error) { - if organizationID != s.organizationID { - return nil, sql.ErrNoRows - } - if s.listCalls.Add(1) == 1 { - return nil, nil - } - return []database.GetEnabledChatModelConfigsByOrganizationRow{{ChatModelConfig: s.config}}, nil -} - -func (s *enabledChatModelConfigsReadSkewStore) GetDefaultChatModelConfig( - _ context.Context, - organizationID uuid.UUID, -) (database.ChatModelConfig, error) { - if organizationID != s.organizationID { - return database.ChatModelConfig{}, sql.ErrNoRows - } - return s.config, nil -} - -type enabledChatModelConfigsEmptyDefaultOrgStore struct { - database.Store - organizationID uuid.UUID -} - -func (s *enabledChatModelConfigsEmptyDefaultOrgStore) GetEnabledChatModelConfigsByOrganization( - _ context.Context, - organizationID uuid.UUID, -) ([]database.GetEnabledChatModelConfigsByOrganizationRow, error) { - if organizationID != s.organizationID { - return nil, xerrors.Errorf("unexpected organization %q", organizationID) - } - return nil, nil -} - -func (s *enabledChatModelConfigsEmptyDefaultOrgStore) GetDefaultChatModelConfig( - _ context.Context, - organizationID uuid.UUID, -) (database.ChatModelConfig, error) { - if organizationID != s.organizationID { - return database.ChatModelConfig{}, xerrors.Errorf("unexpected organization %q", organizationID) - } - return database.ChatModelConfig{}, sql.ErrNoRows -} - -func (s *enabledChatModelConfigsEmptyDefaultOrgStore) GetDefaultOrganization(context.Context) (database.Organization, error) { - return database.Organization{ID: s.organizationID}, nil -} - -func TestEnabledChatModelConfigsWithDefaultOrgFallback(t *testing.T) { - t.Parallel() - - db, _ := dbtestutil.NewDB(t) - ctx := testutil.Context(t, testutil.WaitShort) - - defaultOrg, err := db.GetDefaultOrganization(ctx) - require.NoError(t, err) - otherOrg := dbgen.Organization(t, db, database.Organization{}) - provider := dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{ - Type: database.AIProviderTypeOpenai, - }, "test-key") - // Every write path promotes a default within the org it writes to, so - // an org that owns configs always owns one. Fixtures that insert - // directly must uphold that: the fallback reads it as the ownership - // marker. - defaultOrgConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, - OrganizationID: defaultOrg.ID, - IsDefault: true, - }) - - t.Run("OrgListEmptyFallsBackToDefaultOrg", func(t *testing.T) { - t.Parallel() - - // A fresh org is guaranteed empty even when the test database - // carries seeded configs in previously created orgs. - emptyOrg := dbgen.Organization(t, db, database.Organization{}) - ctx := testutil.Context(t, testutil.WaitShort) - rows, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, db, emptyOrg.ID) - require.NoError(t, err) - require.NotEmpty(t, rows) - - found := slices.ContainsFunc(rows, func(row database.GetEnabledChatModelConfigsByOrganizationRow) bool { - return row.ChatModelConfig.ID == defaultOrgConfig.ID - }) - require.True(t, found, "default org list should include its config") - }) - - t.Run("ReFetchesWhenDefaultAppearsAfterEmptyList", func(t *testing.T) { - t.Parallel() - - organizationID := uuid.New() - config := database.ChatModelConfig{ - ID: uuid.New(), - OrganizationID: organizationID, - IsDefault: true, - } - store := &enabledChatModelConfigsReadSkewStore{ - organizationID: organizationID, - config: config, - } - - rows, err := enabledChatModelConfigsWithDefaultOrgFallback(t.Context(), store, organizationID) - require.NoError(t, err) - require.Len(t, rows, 1) - require.Equal(t, config.ID, rows[0].ChatModelConfig.ID) - require.EqualValues(t, 2, store.listCalls.Load()) - }) - - t.Run("OrgListPresentNeverFallsBack", func(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - - // Give the other org its own enabled config: the default org's - // list must not leak in. - ownConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, - OrganizationID: otherOrg.ID, - }) - - rows, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, db, otherOrg.ID) - require.NoError(t, err) - require.Len(t, rows, 1) - require.Equal(t, ownConfig.ID, rows[0].ChatModelConfig.ID) - }) - - t.Run("DefaultOrgNeverFallsBack", func(t *testing.T) { - t.Parallel() - - store := &enabledChatModelConfigsEmptyDefaultOrgStore{ - Store: db, - organizationID: defaultOrg.ID, - } - got, err := enabledChatModelConfigsWithDefaultOrgFallback(t.Context(), store, defaultOrg.ID) - require.NoError(t, err) - require.Empty(t, got) - }) - - t.Run("OrgWithEveryConfigDisabledNeverFallsBack", func(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - - // The org owns a config but has disabled it, so its enabled - // list is empty. It must keep that empty list instead of - // borrowing the default org's models. - disabledOrg := dbgen.Organization(t, db, database.Organization{}) - dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, - OrganizationID: disabledOrg.ID, - IsDefault: true, - }, func(params *database.InsertChatModelConfigParams) { - // dbgen defaults Enabled to true, so disable it here. - params.Enabled = false - }) - - got, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, db, disabledOrg.ID) - require.NoError(t, err) - require.Empty(t, got) - }) - - t.Run("OrgWithDisabledProviderNeverFallsBack", func(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - - // The org's only config is enabled but its provider is - // disabled, which also yields an empty enabled list. - disabledProviderOrg := dbgen.Organization(t, db, database.Organization{}) - disabledProvider := dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{ - Type: database.AIProviderTypeOpenai, - }, "test-key", func(params *database.InsertAIProviderParams) { - // dbgen defaults Enabled to true, so disable it here. - params.Enabled = false - }) - dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: disabledProvider.ID, Valid: true}, - OrganizationID: disabledProviderOrg.ID, - IsDefault: true, - }) - - got, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, db, disabledProviderOrg.ID) - require.NoError(t, err) - require.Empty(t, got) - }) -} - -func TestListSubagentModels_NonDefaultOrgListIsOrgLocal(t *testing.T) { +func TestListSubagentModels_NonDefaultOrgSeesOnlyOwnOrgConfigs(t *testing.T) { t.Parallel() db, ps := dbtestutil.NewDB(t) @@ -5068,8 +4904,7 @@ func TestListSubagentModels_NonDefaultOrgListIsOrgLocal(t *testing.T) { // The chat's org has its own config (the seed model), so the // list is org-local and the default org's config must not leak - // in. The empty-org fallback is covered by - // TestEnabledChatModelConfigsWithDefaultOrgFallback. + // in. defaultOrg, err := db.GetDefaultOrganization(ctx) require.NoError(t, err) defaultOrgProvider := dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{ diff --git a/coderd/x/chatd/title_override.go b/coderd/x/chatd/title_override.go index 4056fdcfe13..bb2a0ddba46 100644 --- a/coderd/x/chatd/title_override.go +++ b/coderd/x/chatd/title_override.go @@ -83,7 +83,7 @@ func (p *Server) resolveTitleGenerationModelOverride( if err != nil { return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, overrideSet, err } - if !overrideSet { + if !overrideSet || modelConfig.OrganizationID != chat.OrganizationID { return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, false, nil } modelConfig = withResolvedReasoningEffort(modelConfig, overrideEffort) diff --git a/coderd/x/chatd/title_override_internal_test.go b/coderd/x/chatd/title_override_internal_test.go index 7df190a4c40..e52a1b7240d 100644 --- a/coderd/x/chatd/title_override_internal_test.go +++ b/coderd/x/chatd/title_override_internal_test.go @@ -401,7 +401,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnset(t *testing.T) { require.Equal(t, preferredConfig, gotConfig) } -func TestResolveManualTitleModel_NonDefaultOrgUsesDefaultOrgConfigs(t *testing.T) { +func TestResolveManualTitleModel_CrossOrgConfigsAreInvisible(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitShort) @@ -409,29 +409,11 @@ func TestResolveManualTitleModel_NonDefaultOrgUsesDefaultOrgConfigs(t *testing.T db := dbmock.NewMockStore(ctrl) logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) chat, _ := titleOverrideTestChatAndMessages(t) - chat.OrganizationID = uuid.New() // non-default org, no configs of its own - defaultOrgID := uuid.New() - providerID := uuid.New() - preferredConfig := database.ChatModelConfig{ - ID: uuid.New(), - AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, - Model: preferredTitleModels[1].model, - Enabled: true, - OrganizationID: defaultOrgID, - } + chat.OrganizationID = uuid.New() db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil) - // The chat's org has no enabled configs; the selector must fall - // back to the default org's list until the org-scoping cutover. db.EXPECT().GetEnabledChatModelConfigsByOrganization(gomock.Any(), chat.OrganizationID).Return(nil, nil) - db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), chat.OrganizationID). - Return(database.ChatModelConfig{}, sql.ErrNoRows) - db.EXPECT().GetDefaultOrganization(gomock.Any()).Return(database.Organization{ID: defaultOrgID}, nil) - db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), defaultOrgID).Return(preferredConfig, nil) - db.EXPECT().GetEnabledChatModelConfigsByOrganization(gomock.Any(), defaultOrgID).Return([]database.GetEnabledChatModelConfigsByOrganizationRow{ - {ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider}, - }, nil) - db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() + db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), chat.OrganizationID).Return(database.ChatModelConfig{}, sql.ErrNoRows) server := titleOverrideTestServer(db, logger) model, gotConfig, err := server.resolveManualTitleModel( @@ -440,9 +422,9 @@ func TestResolveManualTitleModel_NonDefaultOrgUsesDefaultOrgConfigs(t *testing.T chat, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) - require.NoError(t, err) - require.NotNil(t, model) - require.Equal(t, preferredConfig, gotConfig) + require.ErrorIs(t, err, ErrNoDefaultChatModelConfig) + require.False(t, model.Valid()) + require.Equal(t, database.ChatModelConfig{}, gotConfig) } func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testing.T) { @@ -561,6 +543,37 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUsable(t *testing.T) require.Equal(t, overrideConfig, gotConfig) } +func TestResolveTitleGenerationModelOverride_CrossOrgFallsBack(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + chat, _ := titleOverrideTestChatAndMessages(t) + chat.OrganizationID = uuid.New() + overrideConfig := titleOverrideModelConfig("gpt-4.1", true) + overrideConfig.OrganizationID = uuid.New() + providerID := uuid.New() + overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + + db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) + db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "test-key", + }}, nil).AnyTimes() + + server := titleOverrideTestServer(db, logger) + modelConfig, model, route, overrideSet, err := server.resolveTitleGenerationModelOverride(ctx, chat, modelBuildOptions{}) + require.NoError(t, err) + require.False(t, overrideSet) + require.Equal(t, database.ChatModelConfig{}, modelConfig) + require.False(t, model.Valid()) + require.Equal(t, aiGatewayModelRoute{}, route) +} + func TestResolveManualTitleModel_TitleGenerationOverrideMissingCredentials(t *testing.T) { t.Parallel() diff --git a/enterprise/coderd/exp_chats_test.go b/enterprise/coderd/exp_chats_test.go index 364b82fc682..3f9b9196963 100644 --- a/enterprise/coderd/exp_chats_test.go +++ b/enterprise/coderd/exp_chats_test.go @@ -14,6 +14,7 @@ import ( "github.com/coder/coder/v2/coderd/aibridgedtest" "github.com/coder/coder/v2/coderd/coderdtest" "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbgen" "github.com/coder/coder/v2/coderd/database/dbtestutil" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/util/ptr" @@ -1079,7 +1080,7 @@ func TestCreateChatNonDefaultOrg(t *testing.T) { ctx := testutil.Context(t, testutil.WaitLong) - client, firstUser := coderdenttest.New(t, &coderdenttest.Options{ + client, db, firstUser := coderdenttest.NewWithDatabase(t, &coderdenttest.Options{ Options: &coderdtest.Options{ DeploymentValues: func() *codersdk.DeploymentValues { v := coderdtest.DeploymentValues(t) @@ -1107,6 +1108,13 @@ func TestCreateChatNonDefaultOrg(t *testing.T) { // Create a second (non-default) org via the API. secondOrg := coderdenttest.CreateOrganization(t, client, coderdenttest.CreateOrganizationOptions{}) + dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, + Model: "gpt-4o-mini", + Enabled: true, + IsDefault: true, + OrganizationID: secondOrg.ID, + }) // Create a member with agents-access in both orgs. memberClientRaw, member := coderdtest.CreateAnotherUser( @@ -1148,7 +1156,7 @@ func TestListChats_OrgAdminOnlySeesOwnChats(t *testing.T) { ctx := testutil.Context(t, testutil.WaitLong) - client, firstUser := coderdenttest.New(t, &coderdenttest.Options{ + client, db, firstUser := coderdenttest.NewWithDatabase(t, &coderdenttest.Options{ Options: &coderdtest.Options{ DeploymentValues: func() *codersdk.DeploymentValues { v := coderdtest.DeploymentValues(t) @@ -1176,6 +1184,13 @@ func TestListChats_OrgAdminOnlySeesOwnChats(t *testing.T) { // Create a second (non-default) org. secondOrg := coderdenttest.CreateOrganization(t, client, coderdenttest.CreateOrganizationOptions{}) + dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, + Model: "gpt-4o-mini", + Enabled: true, + IsDefault: true, + OrganizationID: secondOrg.ID, + }) // Create a member with agents-access in both orgs. memberClientRaw, _ := coderdtest.CreateAnotherUser( From 2d922f561b2aff1a195a538f789651a1779e7591 Mon Sep 17 00:00:00 2001 From: Ethan Dickson Date: Mon, 10 Aug 2026 13:06:09 +0000 Subject: [PATCH 3/3] refactor: keep legacy chat model configs in default organization --- ...6_chat_model_config_org_explosion.down.sql | 53 -- ...566_chat_model_config_org_explosion.up.sql | 116 ----- coderd/database/migrations/migrate_test.go | 490 ------------------ ...566_chat_model_config_org_explosion.up.sql | 33 -- coderd/exp_chats.go | 48 +- coderd/exp_chats_test.go | 181 +------ coderd/x/chatd/chatd.go | 119 ++++- coderd/x/chatd/chatd_internal_test.go | 65 ++- coderd/x/chatd/compaction_override.go | 3 - .../compaction_override_internal_test.go | 28 - coderd/x/chatd/configcache.go | 13 +- coderd/x/chatd/configcache_internal_test.go | 71 ++- coderd/x/chatd/subagent.go | 12 +- coderd/x/chatd/subagent_internal_test.go | 237 +++++++-- coderd/x/chatd/title_override.go | 2 +- .../x/chatd/title_override_internal_test.go | 61 +-- enterprise/coderd/exp_chats_test.go | 19 +- 17 files changed, 456 insertions(+), 1095 deletions(-) delete mode 100644 coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql delete mode 100644 coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql delete mode 100644 coderd/database/migrations/testdata/fixtures/000566_chat_model_config_org_explosion.up.sql diff --git a/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql b/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql deleted file mode 100644 index da93da6ec57..00000000000 --- a/coderd/database/migrations/000566_chat_model_config_org_explosion.down.sql +++ /dev/null @@ -1,53 +0,0 @@ --- This migration is best-effort because copied configs have no persisted --- provenance. A non-default config is treated as a copy when a default-org --- config has the same provider and model. -CREATE TEMPORARY TABLE model_config_copy_map ( - copy_id uuid PRIMARY KEY, - orig_id uuid NOT NULL -) ON COMMIT DROP; - -INSERT INTO model_config_copy_map (copy_id, orig_id) -SELECT cp.id, orig.id -FROM chat_model_configs cp -JOIN organizations copy_org - ON copy_org.id = cp.organization_id - AND NOT copy_org.is_default -JOIN LATERAL ( - SELECT default_config.id - FROM chat_model_configs default_config - JOIN organizations default_org - ON default_org.id = default_config.organization_id - AND default_org.is_default - WHERE default_config.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id - AND default_config.model = cp.model - ORDER BY default_config.created_at ASC, default_config.id ASC - LIMIT 1 -) orig ON true; - -UPDATE chats c -SET last_model_config_id = m.orig_id -FROM model_config_copy_map m -WHERE c.last_model_config_id = m.copy_id; - -UPDATE chat_messages mm -SET model_config_id = m.orig_id -FROM model_config_copy_map m -WHERE mm.model_config_id = m.copy_id; - -UPDATE chat_queued_messages q -SET model_config_id = m.orig_id -FROM model_config_copy_map m -WHERE q.model_config_id = m.copy_id; - -UPDATE chat_debug_runs d -SET model_config_id = m.orig_id -FROM model_config_copy_map m -WHERE d.model_config_id = m.copy_id; - -DELETE FROM user_configs uc -USING model_config_copy_map m -WHERE uc.key = 'chat_compaction_threshold_pct:' || m.copy_id::text; - -DELETE FROM chat_model_configs cmc -USING model_config_copy_map m -WHERE cmc.id = m.copy_id; diff --git a/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql b/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql deleted file mode 100644 index ee0ed461d69..00000000000 --- a/coderd/database/migrations/000566_chat_model_config_org_explosion.up.sql +++ /dev/null @@ -1,116 +0,0 @@ --- Copy default-organization chat model configs into each live non-default --- organization. Referenced soft-deleted configs are copied only into the --- organizations that reference them. -CREATE TEMPORARY TABLE model_config_copy_map ( - orig_id uuid NOT NULL, - org_id uuid NOT NULL, - copy_id uuid NOT NULL, - PRIMARY KEY (orig_id, org_id) -) ON COMMIT DROP; - -INSERT INTO model_config_copy_map (orig_id, org_id, copy_id) -SELECT cmc.id, o.id, gen_random_uuid() -FROM chat_model_configs cmc -JOIN organizations def ON def.id = cmc.organization_id AND def.is_default -CROSS JOIN organizations o -WHERE NOT o.is_default - AND NOT o.deleted - AND ( - NOT cmc.deleted - OR EXISTS ( - SELECT 1 - FROM chats c - WHERE c.last_model_config_id = cmc.id - AND c.organization_id = o.id - ) - OR EXISTS ( - SELECT 1 - FROM chat_messages mm - JOIN chats c ON c.id = mm.chat_id - WHERE mm.model_config_id = cmc.id - AND c.organization_id = o.id - ) - OR EXISTS ( - SELECT 1 - FROM chat_queued_messages q - JOIN chats c ON c.id = q.chat_id - WHERE q.model_config_id = cmc.id - AND c.organization_id = o.id - ) - OR EXISTS ( - SELECT 1 - FROM chat_debug_runs d - JOIN chats c ON c.id = d.chat_id - WHERE d.model_config_id = cmc.id - AND c.organization_id = o.id - ) - ); - --- Each copy retains the original behavior and audit fields. The everyone --- group ACL is re-keyed to the destination organization. -INSERT INTO chat_model_configs - (id, model, display_name, created_by, updated_by, enabled, is_default, - deleted, deleted_at, created_at, updated_at, context_limit, - compression_threshold, options, ai_provider_id, organization_id, - group_acl, user_acl) -SELECT - m.copy_id, - cmc.model, - cmc.display_name, - cmc.created_by, - cmc.updated_by, - cmc.enabled, - cmc.is_default, - cmc.deleted, - cmc.deleted_at, - cmc.created_at, - cmc.updated_at, - cmc.context_limit, - cmc.compression_threshold, - cmc.options, - cmc.ai_provider_id, - m.org_id, - jsonb_build_object( - m.org_id::text, - COALESCE( - cmc.group_acl -> cmc.organization_id::text, - '{"permissions": ["read"]}'::jsonb - ) - ), - '{}'::jsonb -FROM model_config_copy_map m -JOIN chat_model_configs cmc ON cmc.id = m.orig_id; - -UPDATE chats c -SET last_model_config_id = m.copy_id -FROM model_config_copy_map m -WHERE c.last_model_config_id = m.orig_id - AND m.org_id = c.organization_id; - -UPDATE chat_messages mm -SET model_config_id = m.copy_id -FROM chats c, model_config_copy_map m -WHERE c.id = mm.chat_id - AND mm.model_config_id = m.orig_id - AND m.org_id = c.organization_id; - -UPDATE chat_queued_messages q -SET model_config_id = m.copy_id -FROM chats c, model_config_copy_map m -WHERE c.id = q.chat_id - AND q.model_config_id = m.orig_id - AND m.org_id = c.organization_id; - -UPDATE chat_debug_runs d -SET model_config_id = m.copy_id -FROM chats c, model_config_copy_map m -WHERE c.id = d.chat_id - AND d.model_config_id = m.orig_id - AND m.org_id = c.organization_id; - -INSERT INTO user_configs (user_id, key, value) -SELECT uc.user_id, 'chat_compaction_threshold_pct:' || m.copy_id::text, uc.value -FROM user_configs uc -JOIN model_config_copy_map m - ON uc.key = 'chat_compaction_threshold_pct:' || m.orig_id::text -ON CONFLICT (user_id, key) DO NOTHING; diff --git a/coderd/database/migrations/migrate_test.go b/coderd/database/migrations/migrate_test.go index 1c13210637f..1f6a1b0b56f 100644 --- a/coderd/database/migrations/migrate_test.go +++ b/coderd/database/migrations/migrate_test.go @@ -2953,493 +2953,3 @@ func mustJSON(t *testing.T, v any) []byte { require.NoError(t, err) return raw } - -func TestMigration000566ChatModelConfigOrgExplosion(t *testing.T) { - t.Parallel() - - const previousMigrationVersion = 565 - - sqlDB := testSQLDB(t) - next, err := migrations.Stepper(sqlDB) - require.NoError(t, err) - for { - version, more, err := next() - require.NoError(t, err) - if !more { - t.Fatalf("migration %d not found", previousMigrationVersion) - } - if version == previousMigrationVersion { - break - } - } - - ctx := testutil.Context(t, testutil.WaitSuperLong) - - now := time.Now().UTC().Truncate(time.Microsecond) - providerID := uuid.New() - user1ID := uuid.New() - user2ID := uuid.New() - orgBID := uuid.New() // live org with chats - orgCID := uuid.New() // live zero-member org: receives the full live set - orgDID := uuid.New() // soft-deleted org: receives nothing, chats untouched - c1ID := uuid.New() // live default config - c2ID := uuid.New() // live plain config - c3ID := uuid.New() // soft-deleted, referenced in orgB only - c4ID := uuid.New() // soft-deleted, unreferenced: never copied - c5ID := uuid.New() // live plain config - emptyACLConfigID := uuid.New() - chatBID := uuid.New() // chat in orgB pinned to live c1 - chatB3ID := uuid.New() // chat in orgB pinned to deleted c3 - chatDID := uuid.New() // chat in soft-deleted orgD pinned to live c1 - - execFixture := func(query string, args ...any) { - t.Helper() - _, err := sqlDB.ExecContext(ctx, query, args...) - require.NoError(t, err) - } - - for i, id := range []uuid.UUID{user1ID, user2ID} { - execFixture( - `INSERT INTO users (id, username, email, hashed_password, created_at, updated_at, status, rbac_roles, login_type) - VALUES ($1, $2, $3, $4, $5, $6, 'active', '{}', 'password')`, - id, fmt.Sprintf("m3user%d", i+1), fmt.Sprintf("m3user%d@coder.com", i+1), []byte{}, now, now, - ) - } - - execFixture( - `INSERT INTO ai_providers (id, type, name, enabled, base_url, created_at, updated_at) - VALUES ($1, $2, $3, $4, $5, $6, $7)`, - providerID, "openai", "openai-566", true, "https://api.openai.com/v1", now, now, - ) - - // Three non-default orgs: live B, live zero-member C, soft-deleted D. - for _, o := range []struct { - id uuid.UUID - name string - deleted bool - }{ - {orgBID, "org-b-566", false}, - {orgCID, "org-c-566", false}, - {orgDID, "org-d-566", true}, - } { - execFixture( - `INSERT INTO organizations (id, name, description, display_name, default_org_member_roles, created_at, updated_at, deleted) - VALUES ($1, $2, '', '', '{}', $3, $3, $4)`, - o.id, o.name, now, o.deleted, - ) - } - - var defaultOrgID uuid.UUID - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT id FROM organizations WHERE is_default = true").Scan(&defaultOrgID)) - - // Model configs in the default org (pre-566 state: every config lives - // there, backfilled by 000565 with the everyone ACL entry). - insertConfig := func(id uuid.UUID, model string, isDefault, deleted bool, acl string) { - t.Helper() - execFixture( - `INSERT INTO chat_model_configs (id, model, display_name, enabled, is_default, deleted, deleted_at, - context_limit, compression_threshold, ai_provider_id, organization_id, group_acl, created_at, updated_at) - VALUES ($1, $2, $3, true, $4, $5, (CASE WHEN $5 THEN $6::timestamptz ELSE NULL END), - 200000, 70, $7, $8, $9::jsonb, $6, $6)`, - id, model, model+" display", isDefault, deleted, now, providerID, defaultOrgID, acl, - ) - } - everyoneACL := `{"` + defaultOrgID.String() + `": {"permissions": ["read"]}}` - insertConfig(c1ID, "gpt-5.2", true, false, everyoneACL) - insertConfig(c2ID, "gpt-5.2-mini", false, false, everyoneACL) - insertConfig(c3ID, "gpt-4-legacy", false, true, everyoneACL) - insertConfig(c4ID, "gpt-4-ancient", false, true, everyoneACL) - insertConfig(c5ID, "gpt-5.2-nano", false, false, everyoneACL) - insertConfig(emptyACLConfigID, "gpt-5.2-empty-acl", false, false, `{}`) - - // Chats: orgB pinned to live c1, orgB second chat pinned to deleted c3, - // soft-deleted orgD pinned to live c1. - for _, ch := range []struct { - id uuid.UUID - orgID uuid.UUID - ownerID uuid.UUID - cfgID uuid.UUID - }{ - {chatBID, orgBID, user1ID, c1ID}, - {chatB3ID, orgBID, user2ID, c3ID}, - {chatDID, orgDID, user1ID, c1ID}, - } { - execFixture( - `INSERT INTO chats (id, owner_id, organization_id, last_model_config_id, created_at, updated_at) - VALUES ($1, $2, $3, $4, $5, $5)`, - ch.id, ch.ownerID, ch.orgID, ch.cfgID, now, - ) - } - - // Messages in each chat referencing the chat's pinned config. - for _, m := range []struct { - chatID uuid.UUID - cfgID uuid.UUID - }{ - {chatBID, c1ID}, - {chatB3ID, c3ID}, - {chatDID, c1ID}, - } { - execFixture( - `INSERT INTO chat_messages (chat_id, model_config_id, role, content, content_version) - VALUES ($1, $2, 'user', '[]'::jsonb, 2)`, - m.chatID, m.cfgID, - ) - } - - // Queued messages (FK-less) referencing configs via their chat's org. - execFixture( - `INSERT INTO chat_queued_messages (chat_id, model_config_id, content, created_by) - VALUES ($1, $2, '[]'::jsonb, $3)`, - chatBID, c1ID, user1ID, - ) - execFixture( - `INSERT INTO chat_queued_messages (chat_id, model_config_id, content, created_by) - VALUES ($1, $2, '[]'::jsonb, $3)`, - chatB3ID, c3ID, user2ID, - ) - execFixture( - `INSERT INTO chat_queued_messages (chat_id, model_config_id, content, created_by) - VALUES ($1, $2, '[]'::jsonb, $3)`, - chatDID, c1ID, user1ID, - ) - - // Debug runs (FK-less, attribution). - execFixture( - `INSERT INTO chat_debug_runs (id, chat_id, model_config_id, kind, status) - VALUES ($1, $2, $3, 'turn', 'finished')`, - uuid.New(), chatBID, c1ID, - ) - execFixture( - `INSERT INTO chat_debug_runs (id, chat_id, model_config_id, kind, status) - VALUES ($1, $2, $3, 'turn', 'finished')`, - uuid.New(), chatB3ID, c3ID, - ) - execFixture( - `INSERT INTO chat_debug_runs (id, chat_id, model_config_id, kind, status) - VALUES ($1, $2, $3, 'turn', 'finished')`, - uuid.New(), chatDID, c1ID, - ) - - // Compaction-threshold keys: c1 (live, users 1+2), c3 (deleted but - // referenced in orgB, user 1), c4 (deleted unreferenced, user 1: must - // never fan out), plus a non-threshold key that must stay untouched. - thresholdKey := func(id uuid.UUID) string { - return "chat_compaction_threshold_pct:" + id.String() - } - for _, tc := range []struct { - userID uuid.UUID - key string - value string - }{ - {user1ID, thresholdKey(c1ID), "80"}, - {user2ID, thresholdKey(c1ID), "75"}, - {user1ID, thresholdKey(c3ID), "60"}, - {user1ID, thresholdKey(c4ID), "55"}, - // Hostile keys: a malformed and an empty suffix. The up leaves - // them alone; the down must not abort on their uuid cast. - {user1ID, "chat_compaction_threshold_pct:not-a-uuid", "50"}, - {user1ID, "chat_compaction_threshold_pct:", "45"}, - {user1ID, "chat_personal_model_override:root", "chat_default"}, - } { - execFixture( - `INSERT INTO user_configs (user_id, key, value) VALUES ($1, $2, $3)`, - tc.userID, tc.key, tc.value, - ) - } - - upSQL, err := os.ReadFile("000566_chat_model_config_org_explosion.up.sql") - require.NoError(t, err) - _, err = sqlDB.ExecContext(ctx, string(upSQL)) - require.NoError(t, err) - - // copyID resolves the copy of orig in org by natural attributes (the - // migration persists no mapping; this is also what the down relies on). - copyID := func(origID, orgID uuid.UUID) (uuid.UUID, bool) { - t.Helper() - var id uuid.UUID - err := sqlDB.QueryRowContext(ctx, - `SELECT cp.id FROM chat_model_configs cp - JOIN chat_model_configs orig ON orig.id = $1 - WHERE cp.organization_id = $2 - AND cp.model = orig.model - AND cp.ai_provider_id IS NOT DISTINCT FROM orig.ai_provider_id - AND cp.id <> orig.id`, origID, orgID).Scan(&id) - if err == sql.ErrNoRows { - return uuid.Nil, false - } - require.NoError(t, err) - return id, true - } - - // --- Per-org row counts --- - // Default org keeps its 6 originals; orgB gets 4 live copies (c1, c2, - // c5, late) plus the referenced-deleted c3 copy; orgC (zero-member) gets - // the full live set (4) and nothing deleted; orgD gets nothing. - assertCount := func(orgID uuid.UUID, want int, msg string) { - t.Helper() - var got int - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT COUNT(*) FROM chat_model_configs WHERE organization_id = $1", orgID).Scan(&got)) - require.Equal(t, want, got, msg) - } - assertCount(defaultOrgID, 6, "default org keeps only its originals") - assertCount(orgBID, 5, "orgB: 4 live fan-out + referenced-deleted c3") - assertCount(orgCID, 4, "orgC: full live set, no deleted copies") - assertCount(orgDID, 0, "soft-deleted org receives no copies") - - var totalConfigs int - require.NoError(t, sqlDB.QueryRowContext(ctx, "SELECT COUNT(*) FROM chat_model_configs").Scan(&totalConfigs)) - require.Equal(t, 15, totalConfigs, "up: 6 originals and 9 copies") - - var duplicateConfigs int - require.NoError(t, sqlDB.QueryRowContext(ctx, ` - SELECT COUNT(*) FROM ( - SELECT organization_id, ai_provider_id, model - FROM chat_model_configs - GROUP BY organization_id, ai_provider_id, model - HAVING COUNT(*) <> 1 - ) duplicates - `).Scan(&duplicateConfigs)) - require.Zero(t, duplicateConfigs, "each organization has one config per provider and model") - - // Copies preserve deleted state and deletion timestamps. - c3CopyB, ok := copyID(c3ID, orgBID) - require.True(t, ok, "orgB received the referenced-deleted c3 copy") - var mismatchedDeletedState int - require.NoError(t, sqlDB.QueryRowContext(ctx, ` - SELECT COUNT(*) - FROM chat_model_configs cp - JOIN chat_model_configs orig - ON orig.organization_id = $1 - AND orig.ai_provider_id IS NOT DISTINCT FROM cp.ai_provider_id - AND orig.model = cp.model - WHERE cp.organization_id IN ($2, $3) - AND (cp.deleted IS DISTINCT FROM orig.deleted - OR cp.deleted_at IS DISTINCT FROM orig.deleted_at) - `, defaultOrgID, orgBID, orgCID).Scan(&mismatchedDeletedState)) - require.Zero(t, mismatchedDeletedState) - - // c4 is copied nowhere. - for _, orgID := range []uuid.UUID{orgBID, orgCID, orgDID} { - _, ok := copyID(c4ID, orgID) - require.False(t, ok, "unreferenced deleted c4 must not be copied") - } - - // --- Exactly one live default per org that received copies --- - for _, orgID := range []uuid.UUID{defaultOrgID, orgBID, orgCID} { - var defaults int - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT COUNT(*) FROM chat_model_configs WHERE organization_id = $1 AND is_default AND NOT deleted", orgID).Scan(&defaults)) - require.Equal(t, 1, defaults, "exactly one live default per org") - } - - // --- Remaps --- - // Chats in live orgB remap to same-org copies; the chat in soft-deleted - // orgD keeps the original reference. - c1CopyB, ok := copyID(c1ID, orgBID) - require.True(t, ok) - assertChatPinned := func(chatID, want uuid.UUID, msg string) { - t.Helper() - var got uuid.UUID - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT last_model_config_id FROM chats WHERE id = $1", chatID).Scan(&got)) - require.Equal(t, want, got, msg) - } - assertChatPinned(chatBID, c1CopyB, "orgB chat remapped to same-org live copy") - assertChatPinned(chatB3ID, c3CopyB, "orgB chat on deleted model remapped to the deleted copy") - assertChatPinned(chatDID, c1ID, "soft-deleted org chat keeps original reference") - - var msgCfg uuid.UUID - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatBID).Scan(&msgCfg)) - require.Equal(t, c1CopyB, msgCfg, "orgB message remapped") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) - require.Equal(t, c3CopyB, msgCfg, "orgB deleted-model message remapped") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatDID).Scan(&msgCfg)) - require.Equal(t, c1ID, msgCfg, "orgD message untouched") - - // Queued messages (FK-less): orgB remapped, orgD untouched. - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatBID).Scan(&msgCfg)) - require.Equal(t, c1CopyB, msgCfg, "orgB queued message remapped") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) - require.Equal(t, c3CopyB, msgCfg, "orgB deleted-model queued message remapped") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatDID).Scan(&msgCfg)) - require.Equal(t, c1ID, msgCfg, "orgD queued message untouched") - - // Debug runs (FK-less): orgB remapped, orgD untouched. - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatBID).Scan(&msgCfg)) - require.Equal(t, c1CopyB, msgCfg, "orgB debug run remapped") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) - require.Equal(t, c3CopyB, msgCfg, "orgB deleted-model debug run remapped") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatDID).Scan(&msgCfg)) - require.Equal(t, c1ID, msgCfg, "orgD debug run untouched") - - // No reference in a live non-default org still points at a default-org - // config. - var dangling int - require.NoError(t, sqlDB.QueryRowContext(ctx, - `SELECT COUNT(*) FROM chats c - JOIN chat_model_configs cmc ON cmc.id = c.last_model_config_id - JOIN organizations def ON def.id = cmc.organization_id AND def.is_default - JOIN organizations co ON co.id = c.organization_id - WHERE NOT co.is_default AND NOT co.deleted`).Scan(&dangling)) - require.Zero(t, dangling, "no live-org chat references a default-org config") - - // Each copy carries only its target organization's everyone entry. - for _, orgID := range []uuid.UUID{orgBID, orgCID} { - rows, err := sqlDB.QueryContext(ctx, - "SELECT group_acl FROM chat_model_configs WHERE organization_id = $1", orgID) - require.NoError(t, err) - for rows.Next() { - var groupACL []byte - require.NoError(t, rows.Scan(&groupACL)) - require.JSONEq(t, string(mustJSON(t, map[string]any{ - orgID.String(): map[string]any{"permissions": []string{"read"}}, - })), string(groupACL)) - } - require.NoError(t, rows.Err()) - require.NoError(t, rows.Close()) - } - // --- Threshold fan-out --- - // c1 keys fanned out to the orgB and orgC copies for both users (follows - // copies, not membership); the c3 key fanned out to the orgB copy only; - // c4 produced nothing. - c1CopyC, ok := copyID(c1ID, orgCID) - require.True(t, ok) - for _, tc := range []struct { - userID uuid.UUID - cfgID uuid.UUID - value string - }{ - {user1ID, c1CopyB, "80"}, - {user1ID, c1CopyC, "80"}, - {user2ID, c1CopyB, "75"}, - {user2ID, c1CopyC, "75"}, - {user1ID, c3CopyB, "60"}, - } { - var value string - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT value FROM user_configs WHERE user_id = $1 AND key = $2", - tc.userID, thresholdKey(tc.cfgID)).Scan(&value)) - require.Equal(t, tc.value, value, "threshold value copied to config %s for user %s", tc.cfgID, tc.userID) - } - var validThresholdCount int - require.NoError(t, sqlDB.QueryRowContext(ctx, ` - SELECT COUNT(*) FROM user_configs - WHERE key ~ '^chat_compaction_threshold_pct:[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' - `).Scan(&validThresholdCount)) - require.Equal(t, 9, validThresholdCount, "4 original and 5 fanned-out threshold keys") - // c3 produced no orgC key (its only copy is in orgB). - var c3CCount int - require.NoError(t, sqlDB.QueryRowContext(ctx, - `SELECT COUNT(*) FROM user_configs uc - JOIN chat_model_configs cp ON cp.id = substring(uc.key FROM 'chat_compaction_threshold_pct:(.*)')::uuid - WHERE uc.key LIKE 'chat_compaction_threshold_pct:%' - AND substring(uc.key FROM 'chat_compaction_threshold_pct:(.*)') ~ '^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' - AND cp.organization_id = $1 AND cp.model = 'gpt-4-legacy'`, - orgCID).Scan(&c3CCount)) - require.Zero(t, c3CCount, "deleted config with no orgC copy produces no orgC key") - // c4 (deleted, unreferenced): the only ancient-model threshold key is the - // seeded original. - var c4Keys []string - rows, err := sqlDB.QueryContext(ctx, - `SELECT key FROM user_configs WHERE key LIKE 'chat_compaction_threshold_pct:%' - AND substring(key FROM 'chat_compaction_threshold_pct:(.*)') ~ '^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' - AND substring(key FROM 'chat_compaction_threshold_pct:(.*)')::uuid = $1`, c4ID) - require.NoError(t, err) - for rows.Next() { - var k string - require.NoError(t, rows.Scan(&k)) - c4Keys = append(c4Keys, k) - } - require.NoError(t, rows.Err()) - require.NoError(t, rows.Close()) - require.Equal(t, []string{thresholdKey(c4ID)}, c4Keys, "unreferenced deleted c4 fans out zero keys") - // Seeded original keys survive; the non-threshold key is untouched. - for _, key := range []string{thresholdKey(c1ID), thresholdKey(c3ID)} { - var exists bool - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT EXISTS(SELECT 1 FROM user_configs WHERE user_id = $1 AND key = $2)", - user1ID, key).Scan(&exists)) - require.True(t, exists, "original threshold key survives") - } - var overrideValue string - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT value FROM user_configs WHERE user_id = $1 AND key = 'chat_personal_model_override:root'", - user1ID).Scan(&overrideValue)) - require.Equal(t, "chat_default", overrideValue, "non-threshold keys are untouched") - - downSQL, err := os.ReadFile("000566_chat_model_config_org_explosion.down.sql") - require.NoError(t, err) - _, err = sqlDB.ExecContext(ctx, string(downSQL)) - require.NoError(t, err) - - // Copies are gone; references restored to the default-org originals. - assertCount(defaultOrgID, 6, "down: default org keeps its originals") - assertCount(orgBID, 0, "down: copies deleted from orgB") - assertCount(orgCID, 0, "down: copies deleted from orgC") - require.NoError(t, sqlDB.QueryRowContext(ctx, "SELECT COUNT(*) FROM chat_model_configs").Scan(&totalConfigs)) - require.Equal(t, 6, totalConfigs, "down: only default-org originals remain") - assertChatPinned(chatBID, c1ID, "down: orgB chat restored to original c1") - assertChatPinned(chatB3ID, c3ID, "down: orgB chat restored to original c3") - assertChatPinned(chatDID, c1ID, "down: orgD chat unchanged") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatBID).Scan(&msgCfg)) - require.Equal(t, c1ID, msgCfg, "down: orgB message restored") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_messages WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) - require.Equal(t, c3ID, msgCfg, "down: orgB deleted-model message restored") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatBID).Scan(&msgCfg)) - require.Equal(t, c1ID, msgCfg, "down: orgB queued message restored") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_queued_messages WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) - require.Equal(t, c3ID, msgCfg, "down: orgB deleted-model queued message restored") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatBID).Scan(&msgCfg)) - require.Equal(t, c1ID, msgCfg, "down: orgB debug run restored") - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT model_config_id FROM chat_debug_runs WHERE chat_id = $1", chatB3ID).Scan(&msgCfg)) - require.Equal(t, c3ID, msgCfg, "down: orgB deleted-model debug run restored") - - // The down removes only fanned-out threshold keys. It preserves all - // seeded keys, including malformed and dangling keys. - var thresholdCount int - require.NoError(t, sqlDB.QueryRowContext(ctx, - "SELECT COUNT(*) FROM user_configs WHERE key LIKE 'chat_compaction_threshold_pct:%'").Scan(&thresholdCount)) - require.Equal(t, 6, thresholdCount, "down: all 6 seeded threshold keys remain") - - _, err = sqlDB.ExecContext(ctx, string(upSQL)) - require.NoError(t, err) - assertCount(defaultOrgID, 6, "re-up: default org keeps its originals") - assertCount(orgBID, 5, "re-up: orgB copies recreated") - assertCount(orgCID, 4, "re-up: orgC copies recreated") - assertCount(orgDID, 0, "re-up: soft-deleted org receives no copies") - require.NoError(t, sqlDB.QueryRowContext(ctx, "SELECT COUNT(*) FROM chat_model_configs").Scan(&totalConfigs)) - require.Equal(t, 15, totalConfigs, "re-up: 6 originals and 9 copies") - require.NoError(t, sqlDB.QueryRowContext(ctx, ` - SELECT COUNT(*) FROM ( - SELECT organization_id, ai_provider_id, model - FROM chat_model_configs - GROUP BY organization_id, ai_provider_id, model - HAVING COUNT(*) <> 1 - ) duplicates - `).Scan(&duplicateConfigs)) - require.Zero(t, duplicateConfigs, "re-up: each organization has one config per provider and model") - require.NoError(t, sqlDB.QueryRowContext(ctx, ` - SELECT COUNT(*) FROM user_configs - WHERE key ~ '^chat_compaction_threshold_pct:[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$' - `).Scan(&validThresholdCount)) - require.Equal(t, 9, validThresholdCount, "re-up: threshold keys fan out once") - assertChatPinned(chatDID, c1ID, "re-up: orgD still untouched") -} diff --git a/coderd/database/migrations/testdata/fixtures/000566_chat_model_config_org_explosion.up.sql b/coderd/database/migrations/testdata/fixtures/000566_chat_model_config_org_explosion.up.sql deleted file mode 100644 index ae2b62dd6c5..00000000000 --- a/coderd/database/migrations/testdata/fixtures/000566_chat_model_config_org_explosion.up.sql +++ /dev/null @@ -1,33 +0,0 @@ --- Fixture for 000566 (org explosion cutover). Fixtures apply right after --- their migration runs, so this executes after the explosion has copied --- the default organization's configs into every live organization. It --- seeds one organically created config in the non-default organization --- from fixture 000291, so later migrations run over configs that the --- explosion did not create. -INSERT INTO chat_model_configs ( - id, - model, - display_name, - enabled, - is_default, - context_limit, - compression_threshold, - ai_provider_id, - organization_id, - group_acl, - created_at, - updated_at -) VALUES ( - '566c0001-0000-4000-8000-000000000001', - 'gpt-5.2-org-fixture', - 'Fixture Org Model 566', - TRUE, - FALSE, - 128000, - 70, - 'a52c6f0e-7d4b-4e1a-9c3f-2b8d5e6f7a8b', - '20362772-802a-4a72-8e4f-3648b4bfd168', - jsonb_build_object('20362772-802a-4a72-8e4f-3648b4bfd168', jsonb_build_object('permissions', jsonb_build_array('read'))), - '2024-01-01 00:00:00+00', - '2024-01-01 00:00:00+00' -); diff --git a/coderd/exp_chats.go b/coderd/exp_chats.go index 794203f692d..ec08fdc96a7 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -1129,7 +1129,6 @@ func (api *API) getUserChatProviderAvailability( func (api *API) userCanUseChatModelConfig( ctx context.Context, userID uuid.UUID, - organizationID uuid.UUID, modelConfigID uuid.UUID, ) (database.ChatModelConfig, chatModelConfigUnavailableReason, error) { if modelConfigID == uuid.Nil { @@ -1146,7 +1145,7 @@ func (api *API) userCanUseChatModelConfig( } return database.ChatModelConfig{}, chatModelConfigAvailable, err } - if model.OrganizationID != organizationID || !model.Enabled { + if !model.Enabled { return database.ChatModelConfig{}, chatModelConfigUnavailableModelNotFoundOrDisabled, nil } @@ -1177,10 +1176,9 @@ func (api *API) userCanUseChatModelConfig( func (api *API) validateUserChatModelConfigAvailable( ctx context.Context, userID uuid.UUID, - organizationID uuid.UUID, modelConfigID uuid.UUID, ) (database.ChatModelConfig, int, *codersdk.Response) { - modelConfig, reason, err := api.userCanUseChatModelConfig(ctx, userID, organizationID, modelConfigID) + modelConfig, reason, err := api.userCanUseChatModelConfig(ctx, userID, modelConfigID) if err != nil { return database.ChatModelConfig{}, http.StatusInternalServerError, &codersdk.Response{ Message: "Internal error validating model config override.", @@ -1221,13 +1219,12 @@ func (api *API) validateUserChatModelConfigAvailable( func (api *API) validateExplicitChatModelConfigAvailable( ctx context.Context, userID uuid.UUID, - organizationID uuid.UUID, modelConfigID uuid.UUID, ) (int, *codersdk.Response) { if modelConfigID == uuid.Nil { return 0, nil } - _, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, organizationID, modelConfigID) + _, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, modelConfigID) return status, resp } @@ -2806,7 +2803,7 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) { if req.ModelConfigID != nil { modelConfigID = *req.ModelConfigID } - if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, chat.OrganizationID, modelConfigID); resp != nil { + if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, modelConfigID); resp != nil { httpapi.Write(ctx, rw, status, *resp) return } @@ -2998,7 +2995,7 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) { if req.ModelConfigID != nil { editModelConfigID = *req.ModelConfigID } - if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, chat.OrganizationID, editModelConfigID); resp != nil { + if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, editModelConfigID); resp != nil { httpapi.Write(ctx, rw, status, *resp) return } @@ -4399,7 +4396,7 @@ func (api *API) resolveCreateChatModelConfigID( Message: "Invalid model config ID.", } } - if _, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, req.OrganizationID, *req.ModelConfigID); resp != nil { + if _, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, *req.ModelConfigID); resp != nil { return uuid.Nil, nil, status, resp } return *req.ModelConfigID, nil, 0, nil @@ -4413,7 +4410,7 @@ func (api *API) resolveCreateChatModelConfigID( } } if !personalOverridesEnabled { - id, status, resp := api.defaultCreateChatModelConfigID(ctx, req.OrganizationID) + id, status, resp := api.defaultCreateChatModelConfigID(ctx) return id, nil, status, resp } @@ -4448,7 +4445,6 @@ func (api *API) resolveCreateChatModelConfigID( _, reason, err := api.userCanUseChatModelConfig( ctx, userID, - req.OrganizationID, parsed.ModelConfigID, ) if err != nil { @@ -4477,15 +4473,26 @@ func (api *API) resolveCreateChatModelConfigID( } } - id, status, resp := api.defaultCreateChatModelConfigID(ctx, req.OrganizationID) + id, status, resp := api.defaultCreateChatModelConfigID(ctx) return id, nil, status, resp } func (api *API) defaultCreateChatModelConfigID( ctx context.Context, - organizationID uuid.UUID, ) (uuid.UUID, int, *codersdk.Response) { - defaultModelConfig, err := api.Database.GetDefaultChatModelConfig(ctx, organizationID) + // The request carries a validated organization, but the pre-cutover + // handler deliberately resolves the deployment-default model until the + // API layer completes the organization-scoping cutover. The default-org + // lookup is internal and does not expose organization data to the user. + //nolint:gocritic // Internal default-org resolution, scoped to this call. + defaultOrg, err := api.Database.GetDefaultOrganization(dbauthz.AsChatd(ctx)) + if err != nil { + return uuid.Nil, http.StatusInternalServerError, &codersdk.Response{ + Message: "Failed to resolve chat model config.", + Detail: err.Error(), + } + } + defaultModelConfig, err := api.Database.GetDefaultChatModelConfig(ctx, defaultOrg.ID) if err != nil { if xerrors.Is(err, sql.ErrNoRows) { return uuid.Nil, http.StatusBadRequest, &codersdk.Response{ @@ -5022,18 +5029,7 @@ func (api *API) putUserChatPersonalModelOverride(rw http.ResponseWriter, r *http }) return } - // Personal model overrides are user-global. They select from the - // default organization's model configs. - //nolint:gocritic // This lookup resolves internal deployment configuration. - defaultOrg, err := api.Database.GetDefaultOrganization(dbauthz.AsChatd(ctx)) - if err != nil { - httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ - Message: "Internal error validating model config override.", - Detail: err.Error(), - }) - return - } - modelConfig, status, resp := api.validateUserChatModelConfigAvailable(ctx, apiKey.UserID, defaultOrg.ID, parsedModelConfigID) + modelConfig, status, resp := api.validateUserChatModelConfigAvailable(ctx, apiKey.UserID, parsedModelConfigID) if resp != nil { httpapi.Write(ctx, rw, status, *resp) return diff --git a/coderd/exp_chats_test.go b/coderd/exp_chats_test.go index 95bc72c9450..d23811a3bb3 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -327,83 +327,32 @@ func insertAssistantMessage( func TestPostChats(t *testing.T) { t.Parallel() - t.Run("SuccessNonDefaultOrgUsesOrgDefault", func(t *testing.T) { + t.Run("SuccessNonDefaultOrgUsesDeploymentDefault", func(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitLong) client, db := newChatClientWithDatabase(t) _ = coderdtest.CreateFirstUser(t, client.Client) - defaultConfig := createChatModelConfig(t, client) + modelConfig := createChatModelConfig(t, client) + + // A member of a non-default org cannot read the default + // organization object, but omitting model_config_id must still + // resolve the deployment default while every config lives there. org := dbgen.Organization(t, db, database.Organization{IsDefault: false}) - orgConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: defaultConfig.AIProviderID, Valid: true}, - Model: "org-default-" + uuid.NewString(), - Enabled: true, - IsDefault: true, - OrganizationID: org.ID, - }) memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, org.ID, rbac.ScopedRoleAgentsAccess(org.ID)) memberClient := codersdk.NewExperimentalClient(memberClientRaw) chat, err := memberClient.CreateChat(ctx, codersdk.CreateChatRequest{ OrganizationID: org.ID, - Content: []codersdk.ChatInputPart{{ - Type: codersdk.ChatInputPartTypeText, - Text: "hello from a non-default org", - }}, + Content: []codersdk.ChatInputPart{ + { + Type: codersdk.ChatInputPartTypeText, + Text: "hello from a non-default org", + }, + }, }) require.NoError(t, err) - require.Equal(t, orgConfig.ID, chat.LastModelConfigID) - }) - - t.Run("NonDefaultOrgWithoutDefaultRejected", func(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitLong) - client, db := newChatClientWithDatabase(t) - _ = coderdtest.CreateFirstUser(t, client.Client) - _ = createChatModelConfig(t, client) - org := dbgen.Organization(t, db, database.Organization{IsDefault: false}) - memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, org.ID, rbac.ScopedRoleAgentsAccess(org.ID)) - memberClient := codersdk.NewExperimentalClient(memberClientRaw) - - _, err := memberClient.CreateChat(ctx, codersdk.CreateChatRequest{ - OrganizationID: org.ID, - Content: []codersdk.ChatInputPart{{ - Type: codersdk.ChatInputPartTypeText, - Text: "no model is configured", - }}, - }) - sdkErr := requireSDKError(t, err, http.StatusBadRequest) - require.Equal(t, "No default chat model config is configured.", sdkErr.Message) - }) - - t.Run("CrossOrgExplicitModelRejected", func(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitLong) - client, db := newChatClientWithDatabase(t) - firstUser := coderdtest.CreateFirstUser(t, client.Client) - defaultConfig := createChatModelConfig(t, client) - otherOrg := dbgen.Organization(t, db, database.Organization{IsDefault: false}) - otherConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: defaultConfig.AIProviderID, Valid: true}, - Model: "cross-org-" + uuid.NewString(), - Enabled: true, - IsDefault: true, - OrganizationID: otherOrg.ID, - }) - - _, err := client.CreateChat(ctx, codersdk.CreateChatRequest{ - OrganizationID: firstUser.OrganizationID, - Content: []codersdk.ChatInputPart{{ - Type: codersdk.ChatInputPartTypeText, - Text: "reject another organization's model", - }}, - ModelConfigID: ptr.Ref(otherConfig.ID), - }) - sdkErr := requireSDKError(t, err, http.StatusBadRequest) - require.Equal(t, "Invalid model_config_id: model config not found or disabled.", sdkErr.Message) + require.Equal(t, modelConfig.ID, chat.LastModelConfigID) }) t.Run("Success", func(t *testing.T) { @@ -7210,40 +7159,6 @@ func TestPostChatMessages(t *testing.T) { require.Equal(t, "Invalid model_config_id: provider is not enabled for this model.", sdkErr.Message) }) - t.Run("CrossOrgModelConfigRejected", func(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitLong) - client, db := newChatClientWithDatabase(t) - firstUser := coderdtest.CreateFirstUser(t, client.Client) - defaultConfig := createChatModelConfig(t, client) - chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{ - OrganizationID: firstUser.OrganizationID, - Content: []codersdk.ChatInputPart{{ - Type: codersdk.ChatInputPartTypeText, - Text: "initial message before cross-org switch", - }}, - }) - require.NoError(t, err) - - otherOrg := dbgen.Organization(t, db, database.Organization{IsDefault: false}) - otherConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: defaultConfig.AIProviderID, Valid: true}, - Model: "cross-org-send-" + uuid.NewString(), - Enabled: true, - OrganizationID: otherOrg.ID, - }) - _, err = client.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{ - Content: []codersdk.ChatInputPart{{ - Type: codersdk.ChatInputPartTypeText, - Text: "reject another organization's model", - }}, - ModelConfigID: ptr.Ref(otherConfig.ID), - }) - sdkErr := requireSDKError(t, err, http.StatusBadRequest) - require.Equal(t, "Invalid model_config_id: model config not found or disabled.", sdkErr.Message) - }) - t.Run("ProviderDisabledDefaultFallbackRejected", func(t *testing.T) { t.Parallel() @@ -8759,47 +8674,6 @@ func TestPatchChatMessage(t *testing.T) { require.False(t, foundOriginalInChat) }) - t.Run("CrossOrgModelConfigRejected", func(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitLong) - client, db := newChatClientWithDatabase(t) - firstUser := coderdtest.CreateFirstUser(t, client.Client) - defaultConfig := createChatModelConfig(t, client) - chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{ - OrganizationID: firstUser.OrganizationID, - Content: []codersdk.ChatInputPart{{ - Type: codersdk.ChatInputPartTypeText, - Text: "before cross-org edit", - }}, - }) - require.NoError(t, err) - messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil) - require.NoError(t, err) - userMessageID := messagesResult.Messages[0].ID - - otherOrg := dbgen.Organization(t, db, database.Organization{IsDefault: false}) - otherConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: defaultConfig.AIProviderID, Valid: true}, - Model: "cross-org-edit-" + uuid.NewString(), - Enabled: true, - OrganizationID: otherOrg.ID, - }) - _, err = client.EditChatMessage(ctx, chat.ID, userMessageID, codersdk.EditChatMessageRequest{ - Content: []codersdk.ChatInputPart{{ - Type: codersdk.ChatInputPartTypeText, - Text: "reject another organization's model", - }}, - ModelConfigID: ptr.Ref(otherConfig.ID), - }) - sdkErr := requireSDKError(t, err, http.StatusBadRequest) - require.Equal(t, "Invalid model_config_id: model config not found or disabled.", sdkErr.Message) - - storedChat, err := db.GetChatByID(dbauthz.AsSystemRestricted(ctx), chat.ID) - require.NoError(t, err) - require.Equal(t, defaultConfig.ID, storedChat.LastModelConfigID) - }) - t.Run("ReasoningEffort", func(t *testing.T) { t.Parallel() @@ -13927,35 +13801,6 @@ func TestCreateChatPersonalModelOverrideRoot(t *testing.T) { require.Equal(t, ptr.Ref("high"), chat.LastReasoningEffort) }) - t.Run("CrossOrgRootModelFallsBackToOrgDefault", func(t *testing.T) { - org := dbgen.Organization(t, db, database.Organization{IsDefault: false}) - orgModel := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: defaultModel.AIProviderID, Valid: true}, - Model: "org-root-personal-" + uuid.NewString(), - Enabled: true, - IsDefault: true, - OrganizationID: org.ID, - }) - otherClientRaw, otherUser := coderdtest.CreateAnotherUser( - t, - adminClient.Client, - org.ID, - rbac.ScopedRoleAgentsAccess(org.ID), - ) - otherClient := codersdk.NewExperimentalClient(otherClientRaw) - upsertRootRaw(otherUser.ID, "model:"+overrideModel.ID.String()) - - chat, err := otherClient.CreateChat(ctx, codersdk.CreateChatRequest{ - OrganizationID: org.ID, - Content: []codersdk.ChatInputPart{{ - Type: codersdk.ChatInputPartTypeText, - Text: "cross-org root model falls back", - }}, - }) - require.NoError(t, err) - require.Equal(t, orgModel.ID, chat.LastModelConfigID) - }) - t.Run("UnavailableRootModelFallsBackToDefault", func(t *testing.T) { upsertRootRaw(firstUser.UserID, "model:"+disabledModel.ID.String()) chat := createChat(adminClient, "disabled root model falls back", nil) diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 4bf3c39dfb4..8334fd67c13 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -1237,7 +1237,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C deploymentPrompt := p.resolveDeploymentSystemPrompt(ctx) if opts.ModelConfigID != uuid.Nil { - if err := requireEnabledChatModelConfig(ctx, p.db, opts.OrganizationID, opts.ModelConfigID); err != nil { + if err := requireEnabledChatModelConfig(ctx, p.db, opts.ModelConfigID); err != nil { return database.Chat{}, err } } @@ -1251,7 +1251,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C contentParts := opts.InitialUserContent if p.hooks.Enabled() { // Validate model admission before dispatch, matching the insert path. - if err := validateCreateModelConfigID(ctx, p.db, opts.OrganizationID, opts.ModelConfigID); err != nil { + if err := validateCreateModelConfigID(ctx, p.db, opts.ModelConfigID); err != nil { return database.Chat{}, err } turnID := uuid.New() @@ -1549,7 +1549,7 @@ func resolveSendMessageModelConfigID( return resolveFallbackModelConfigID(ctx, store, chat.OrganizationID, chat.LastModelConfigID) } - if err := requireEnabledChatModelConfig(ctx, store, chat.OrganizationID, requested); err != nil { + if err := requireEnabledChatModelConfig(ctx, store, requested); err != nil { return uuid.Nil, err } return requested, nil @@ -1560,12 +1560,10 @@ func resolveSendMessageModelConfigID( func requireEnabledChatModelConfig( ctx context.Context, store database.Store, - organizationID uuid.UUID, modelConfigID uuid.UUID, ) error { chatdCtx := chatdModelConfigLookupContext(ctx) - modelConfig, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID) - if err != nil { + if _, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID); err != nil { if errors.Is(err, sql.ErrNoRows) { return xerrors.Errorf( "%w: %s", @@ -1579,27 +1577,20 @@ func requireEnabledChatModelConfig( err, ) } - if modelConfig.OrganizationID != organizationID { - return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID) - } return nil } -func validateCreateModelConfigID(ctx context.Context, store database.Store, organizationID, modelConfigID uuid.UUID) error { +func validateCreateModelConfigID(ctx context.Context, store database.Store, modelConfigID uuid.UUID) error { if modelConfigID == uuid.Nil { return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID) } chatdCtx := chatdModelConfigLookupContext(ctx) - modelConfig, err := store.GetChatModelConfigByID(chatdCtx, modelConfigID) - if err != nil { + if _, err := store.GetChatModelConfigByID(chatdCtx, modelConfigID); err != nil { if errors.Is(err, sql.ErrNoRows) { return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID) } return xerrors.Errorf("get requested model config %s: %w", modelConfigID, err) } - if modelConfig.OrganizationID != organizationID { - return xerrors.Errorf("%w: %s", ErrInvalidModelConfigID, modelConfigID) - } return nil } @@ -1611,10 +1602,8 @@ func resolveFallbackModelConfigID( ) (uuid.UUID, error) { chatdCtx := chatdModelConfigLookupContext(ctx) if modelConfigID != uuid.Nil { - if modelConfig, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID); err == nil { - if modelConfig.OrganizationID == organizationID { - return modelConfigID, nil - } + if _, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID); err == nil { + return modelConfigID, nil } else if !errors.Is(err, sql.ErrNoRows) { return uuid.Nil, xerrors.Errorf( "get chat model config %s: %w", @@ -1624,7 +1613,7 @@ func resolveFallbackModelConfigID( } } - defaultConfig, err := store.GetDefaultChatModelConfig(chatdCtx, organizationID) + defaultConfig, err := defaultChatModelConfigForOrg(chatdCtx, store, organizationID) if err != nil { if errors.Is(err, sql.ErrNoRows) { return uuid.Nil, ErrNoDefaultChatModelConfig @@ -1652,13 +1641,12 @@ func resolveFallbackModelConfigID( func validateModelConfigOverride( ctx context.Context, store database.Store, - organizationID uuid.UUID, requested uuid.UUID, ) (uuid.NullUUID, error) { if requested == uuid.Nil { return uuid.NullUUID{}, nil } - if err := requireEnabledChatModelConfig(ctx, store, organizationID, requested); err != nil { + if err := requireEnabledChatModelConfig(ctx, store, requested); err != nil { return uuid.NullUUID{}, err } return uuid.NullUUID{UUID: requested, Valid: true}, nil @@ -1681,6 +1669,87 @@ func validateEditTarget(ctx context.Context, store database.Store, chatID uuid.U return nil } +// defaultChatModelConfigForOrg resolves the default model config that +// serves organizationID. An org that owns no configs of its own reads +// the default org's default instead, which preserves the +// pre-org-scoping behavior where every chat saw the deployment-wide +// configs. The default org never falls back: a miss there is a real +// absence and returns sql.ErrNoRows. +// +// The returned config's OrganizationID identifies which org's configs +// apply, so callers that need the whole list can read it from there. +// TODO(mafredri): remove after CODAGT-709 M3 (org-scoping cutover); +// per-org resolution becomes strict. +func defaultChatModelConfigForOrg( + ctx context.Context, + store database.Store, + organizationID uuid.UUID, +) (database.ChatModelConfig, error) { + config, err := store.GetDefaultChatModelConfig(ctx, organizationID) + if err == nil { + return config, nil + } + if !errors.Is(err, sql.ErrNoRows) { + return database.ChatModelConfig{}, err + } + defaultOrg, err := store.GetDefaultOrganization(ctx) + if err != nil { + return database.ChatModelConfig{}, xerrors.Errorf("get default organization: %w", err) + } + if defaultOrg.ID == organizationID { + return database.ChatModelConfig{}, sql.ErrNoRows + } + return store.GetDefaultChatModelConfig(ctx, defaultOrg.ID) +} + +// enabledChatModelConfigsWithDefaultOrgFallback returns the organization's +// enabled configs. It uses the default org's configs when the organization +// owns no configs, which preserves the pre-org-scoping behavior. The default +// org itself never falls back. +// +// An empty initial result does not prove the organization owns no configs. Its +// configs can all be disabled or use disabled providers. The resolved default +// config is the ownership marker because every write path preserves one default +// in each organization that owns configs. +// TODO(mafredri): remove after CODAGT-709 M3 (org-scoping cutover); +// organizations list strictly within their own configs. +func enabledChatModelConfigsWithDefaultOrgFallback( + ctx context.Context, + store database.Store, + organizationID uuid.UUID, +) ([]database.GetEnabledChatModelConfigsByOrganizationRow, error) { + rows, err := store.GetEnabledChatModelConfigsByOrganization(ctx, organizationID) + if err != nil { + return nil, err + } + if len(rows) > 0 { + return rows, nil + } + + defaultConfig, err := defaultChatModelConfigForOrg(ctx, store, organizationID) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return rows, nil + } + return nil, xerrors.Errorf("resolve default chat model config: %w", err) + } + if defaultConfig.OrganizationID == organizationID { + // The default can appear after the initial list read. Re-fetch the + // organization's enabled configs to avoid returning a stale empty list. + rows, err = store.GetEnabledChatModelConfigsByOrganization(ctx, organizationID) + if err != nil { + return nil, xerrors.Errorf("re-fetch organization enabled chat model configs: %w", err) + } + return rows, nil + } + + fallbackRows, err := store.GetEnabledChatModelConfigsByOrganization(ctx, defaultConfig.OrganizationID) + if err != nil { + return nil, xerrors.Errorf("get default org enabled chat model configs: %w", err) + } + return fallbackRows, nil +} + // EditMessage replaces an earlier user message and discards the // active-history suffix through chatstate.EditMessage. Model-config // override validation and usage-limit admission run in the same @@ -1714,7 +1783,7 @@ func (p *Server) EditMessage( if err := validateEditTarget(ctx, p.db, opts.ChatID, opts.EditedMessageID); err != nil { return EditMessageResult{}, err } - if _, err := validateModelConfigOverride(ctx, p.db, chat.OrganizationID, opts.ModelConfigID); err != nil { + if _, err := validateModelConfigOverride(ctx, p.db, opts.ModelConfigID); err != nil { return EditMessageResult{}, err } sessionStartHookResult, err = p.hooks.Trigger(ctx, chathooks.ChatFor(chat, &turnID), chathooks.Message{Source: chathooks.SessionStartSourceClear}, agenthooks.EventSessionStart, dispatch.CapacityClassAdmission) @@ -1774,7 +1843,7 @@ func (p *Server) EditMessage( } editedMsg = target - modelOverride, err := validateModelConfigOverride(ctx, store, lockedChat.OrganizationID, opts.ModelConfigID) + modelOverride, err := validateModelConfigOverride(ctx, store, opts.ModelConfigID) if err != nil { return err } @@ -2760,7 +2829,7 @@ func (p *Server) resolveManualTitleModel( return overrideModel, overrideConfig, nil } - configs, err := store.GetEnabledChatModelConfigsByOrganization(ctx, chat.OrganizationID) + configs, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, store, chat.OrganizationID) if err != nil { p.logger.Debug(ctx, "failed to list manual title model configs", slog.F("chat_id", chat.ID), diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index 1d5ae1bf3bc..24bd6b2e45d 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -1129,9 +1129,13 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) { }, ).Return(nil, nil) db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil) - // An empty org list falls through to the chat's fallback model; - // strict org scoping reads no other org's configs. db.EXPECT().GetEnabledChatModelConfigsByOrganization(gomock.Any(), gomock.Any()).Return(nil, nil) + // An empty org list only triggers the pre-cutover fallback when the + // org owns no configs at all, which the missing default proves. The + // default org has no default either, so the empty list stands. + db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), gomock.Any()). + Return(database.ChatModelConfig{}, sql.ErrNoRows).Times(2) + db.EXPECT().GetDefaultOrganization(gomock.Any()).Return(database.Organization{ID: uuid.New()}, nil) db.EXPECT().InTx(gomock.Any(), nil).DoAndReturn( func(fn func(database.Store) error, opts *database.TxOptions) error { @@ -1276,9 +1280,13 @@ func TestRegenerateChatTitle_SkipsPersistWhenTitleChangedConcurrently(t *testing }, ).Return(nil, nil) db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil) - // An empty org list falls through to the chat's fallback model; - // strict org scoping reads no other org's configs. db.EXPECT().GetEnabledChatModelConfigsByOrganization(gomock.Any(), gomock.Any()).Return(nil, nil) + // An empty org list only triggers the pre-cutover fallback when the + // org owns no configs at all, which the missing default proves. The + // default org has no default either, so the empty list stands. + db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), gomock.Any()). + Return(database.ChatModelConfig{}, sql.ErrNoRows).Times(2) + db.EXPECT().GetDefaultOrganization(gomock.Any()).Return(database.Organization{ID: uuid.New()}, nil) db.EXPECT().InTx(gomock.Any(), nil).DoAndReturn( func(fn func(database.Store) error, _ *database.TxOptions) error { @@ -3797,38 +3805,41 @@ func TestResolveFallbackModelConfigID(t *testing.T) { require.Equal(t, defaultModel.ID, resolved) }) - t.Run("NonDefaultOrgWithoutOwnDefaultMisses", func(t *testing.T) { + t.Run("NonDefaultOrgFallsBackToDefaultOrgDefault", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) ctx := testutil.Context(t, testutil.WaitShort) - // The chat's org has no configs of its own; the default org has - // one. Strict scoping resolves configs only within the chat's - // org, so the lookup reports no default. + // The chat's org has no configs of its own; the deployment + // default lives in the default org. Pre-cutover behavior must + // resolve it for chats in any org. otherOrgID := newModelConfigOrg(t, db) defaultOrg, err := db.GetDefaultOrganization(ctx) require.NoError(t, err) provider := newProvider(t, db, true) - _ = newModelConfig(t, db, defaultOrg.ID, provider.ID, true) + defaultModel := newModelConfig(t, db, defaultOrg.ID, provider.ID, true) - _, err = resolveFallbackModelConfigID(ctx, db, otherOrgID, uuid.Nil) - require.ErrorIs(t, err, ErrNoDefaultChatModelConfig) + resolved, err := resolveFallbackModelConfigID(ctx, db, otherOrgID, uuid.Nil) + require.NoError(t, err) + require.Equal(t, defaultModel.ID, resolved) }) - t.Run("OrgDefaultResolvesAsChatd", func(t *testing.T) { + t.Run("FallbackReadsDefaultOrgAsChatd", func(t *testing.T) { t.Parallel() db, _ := dbtestutil.NewDB(t) ctx := testutil.Context(t, testutil.WaitShort) - // The fallback reads the org default under the chatd subject, - // which must be authorized to read chat model configs or every - // fallback path fails closed. - orgID := newModelConfigOrg(t, db) + // The pre-cutover fallback reads the default organization under + // the chatd subject, which must be authorized to read + // organizations or every fallback path fails closed. + otherOrgID := newModelConfigOrg(t, db) + defaultOrg, err := db.GetDefaultOrganization(ctx) + require.NoError(t, err) provider := newProvider(t, db, true) - defaultModel := newModelConfig(t, db, orgID, provider.ID, true) + defaultModel := newModelConfig(t, db, defaultOrg.ID, provider.ID, true) chatdCtx := dbauthz.AsChatd(ctx) - resolved, err := resolveFallbackModelConfigID(chatdCtx, db, orgID, uuid.Nil) + resolved, err := resolveFallbackModelConfigID(chatdCtx, db, otherOrgID, uuid.Nil) require.NoError(t, err) require.Equal(t, defaultModel.ID, resolved) }) @@ -3856,25 +3867,11 @@ func TestResolveFallbackModelConfigID(t *testing.T) { provider := newProvider(t, db, true) model := newModelConfig(t, db, orgID, provider.ID, false) - resolved, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{OrganizationID: orgID}, model.ID) + resolved, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{}, model.ID) require.NoError(t, err) require.Equal(t, model.ID, resolved) }) - t.Run("ExplicitCrossOrgModelRejected", func(t *testing.T) { - t.Parallel() - db, _ := dbtestutil.NewDB(t) - ctx := testutil.Context(t, testutil.WaitShort) - - chatOrgID := newModelConfigOrg(t, db) - modelOrgID := newModelConfigOrg(t, db) - provider := newProvider(t, db, true) - model := newModelConfig(t, db, modelOrgID, provider.ID, false) - - _, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{OrganizationID: chatOrgID}, model.ID) - require.ErrorIs(t, err, ErrInvalidModelConfigID) - }) - // An explicit model whose provider was disabled after the coderd // preflight must still be rejected inside the daemon. t.Run("ExplicitProviderDisabledRejected", func(t *testing.T) { @@ -3886,7 +3883,7 @@ func TestResolveFallbackModelConfigID(t *testing.T) { disabledProvider := newProvider(t, db, false) model := newModelConfig(t, db, orgID, disabledProvider.ID, false) - _, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{OrganizationID: orgID}, model.ID) + _, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{}, model.ID) require.ErrorIs(t, err, ErrInvalidModelConfigID) }) diff --git a/coderd/x/chatd/compaction_override.go b/coderd/x/chatd/compaction_override.go index 8068da64613..fc764ab58db 100644 --- a/coderd/x/chatd/compaction_override.go +++ b/coderd/x/chatd/compaction_override.go @@ -89,9 +89,6 @@ func (p *Server) resolveCompactionOverrideConfig( if err != nil || !overrideSet { return nil, err } - if modelConfig.OrganizationID != chat.OrganizationID { - return nil, err - } // Already validated by the shared resolver; failure is unreachable. resolvedProvider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint( modelConfig.Model, diff --git a/coderd/x/chatd/compaction_override_internal_test.go b/coderd/x/chatd/compaction_override_internal_test.go index 24a94f4b0f9..166263c3d52 100644 --- a/coderd/x/chatd/compaction_override_internal_test.go +++ b/coderd/x/chatd/compaction_override_internal_test.go @@ -146,34 +146,6 @@ func TestResolveCompactionOverrideConfig_DisabledConfigFallsBack(t *testing.T) { require.Nil(t, override) } -func TestResolveCompactionOverrideConfig_CrossOrgFallsBack(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - ctrl := gomock.NewController(t) - db := dbmock.NewMockStore(ctrl) - logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) - chat, _ := titleOverrideTestChatAndMessages(t) - chat.OrganizationID = uuid.New() - overrideConfig := titleOverrideModelConfig("gpt-4.1", true) - overrideConfig.OrganizationID = uuid.New() - providerID := uuid.New() - overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} - - db.EXPECT().GetChatCompactionModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) - db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) - db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() - db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ - ProviderID: providerID, - APIKey: "test-key", - }}, nil).AnyTimes() - - server := titleOverrideTestServer(db, logger) - override, err := server.resolveCompactionOverrideConfig(ctx, chat) - require.NoError(t, err) - require.Nil(t, override) -} - func TestResolveCompactionOverrideConfig_MissingCredentialsFallsBack(t *testing.T) { t.Parallel() diff --git a/coderd/x/chatd/configcache.go b/coderd/x/chatd/configcache.go index a1e9cb34fef..213905e1be6 100644 --- a/coderd/x/chatd/configcache.go +++ b/coderd/x/chatd/configcache.go @@ -286,8 +286,11 @@ func (c *chatConfigCache) storeModelConfig(snap modelConfigSnapshot, config data } // DefaultModelConfig returns the default model config for the given -// organization. Orgs resolve strictly within their own configs; an org -// without its own default reports absence. +// organization. Until the org-scoping cutover, an org without its own +// default falls back to the default org's config. +// TODO(mafredri): remove after CODAGT-709 M3 (org-scoping cutover); +// orgs resolve strictly within their own configs after the org-scoping +// cutover. func (c *chatConfigCache) DefaultModelConfig(ctx context.Context, orgID uuid.UUID) (database.ChatModelConfig, error) { if config, ok := c.cachedDefaultModelConfig(orgID); ok { return config, nil @@ -299,7 +302,7 @@ func (c *chatConfigCache) DefaultModelConfig(ctx context.Context, orgID uuid.UUI return cached, nil } - fetched, err := c.db.GetDefaultChatModelConfig(c.ctx, orgID) + fetched, err := defaultChatModelConfigForOrg(c.ctx, c.db, orgID) if err != nil { return database.ChatModelConfig{}, err } @@ -420,8 +423,8 @@ func (c *chatConfigCache) InvalidateModelConfig(id uuid.UUID) { delete(c.modelConfigs, id) c.modelTopologyEpoch++ // Coarse invalidation: the event does not say whether the changed - // config was a default, nor for which org, so every per-org default - // is dropped. + // config was a default, nor which orgs resolve to it through the + // default-org fallback, so every per-org default is dropped. clear(c.defaultModelConfigs) c.defaultModelConfigGeneration++ c.mu.Unlock() diff --git a/coderd/x/chatd/configcache_internal_test.go b/coderd/x/chatd/configcache_internal_test.go index 9b53787b4f3..92da3c2bbf2 100644 --- a/coderd/x/chatd/configcache_internal_test.go +++ b/coderd/x/chatd/configcache_internal_test.go @@ -336,44 +336,91 @@ func TestConfigCache_DefaultModelConfig_PerOrgKeying(t *testing.T) { require.Equal(t, int32(2), store.defaultModelConfigCallCount(orgB)) } -func TestConfigCache_DefaultModelConfig_CrossOrgIsolation(t *testing.T) { +func TestConfigCache_DefaultModelConfig_DefaultOrgFallback(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitShort) clock := quartz.NewMock(t) + defaultOrgID := uuid.New() otherOrgID := uuid.New() defaultOrgConfig := testChatModelConfig(uuid.New(), "default-org-model") store := &stubChatConfigStore{} store.getDefaultChatModelConfig = func(_ context.Context, orgID uuid.UUID) (database.ChatModelConfig, error) { - // The chat's org resolves no default; the default org has one. + if orgID == defaultOrgID { + return defaultOrgConfig, nil + } return database.ChatModelConfig{}, sql.ErrNoRows } + store.getDefaultOrganization = func(context.Context) (database.Organization, error) { + return database.Organization{ID: defaultOrgID}, nil + } cache := newChatConfigCache(ctx, store, clock) - // Orgs resolve strictly within their own configs: an org without its - // own default reports absence and never reads another org's default. - _, err := cache.DefaultModelConfig(ctx, otherOrgID) - require.ErrorIs(t, err, sql.ErrNoRows) + // An org without its own default resolves the default org's config, + // cached under its own org key. + resolved, err := cache.DefaultModelConfig(ctx, otherOrgID) + require.NoError(t, err) + require.Equal(t, defaultOrgConfig, resolved) + resolvedAgain, err := cache.DefaultModelConfig(ctx, otherOrgID) + require.NoError(t, err) + require.Equal(t, defaultOrgConfig, resolvedAgain) require.Equal(t, int32(1), store.defaultModelConfigCallCount(otherOrgID)) - require.Equal(t, int32(0), store.defaultModelConfigCallCount(defaultOrgConfig.OrganizationID)) + require.Equal(t, int32(1), store.defaultModelConfigCallCount(defaultOrgID)) + require.Equal(t, int32(1), store.defaultOrganizationCall.Load()) + + // The default org itself gets no fallback. + _, err = cache.DefaultModelConfig(ctx, defaultOrgID) + require.NoError(t, err) + require.Equal(t, int32(2), store.defaultModelConfigCallCount(defaultOrgID)) + require.Equal(t, int32(1), store.defaultOrganizationCall.Load()) } -func TestConfigCache_DefaultModelConfig_Miss(t *testing.T) { +func TestConfigCache_DefaultModelConfig_DefaultOrgMiss(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitShort) clock := quartz.NewMock(t) - orgID := uuid.New() + defaultOrgID := uuid.New() + store := &stubChatConfigStore{} + store.getDefaultChatModelConfig = func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { + return database.ChatModelConfig{}, sql.ErrNoRows + } + store.getDefaultOrganization = func(context.Context) (database.Organization, error) { + return database.Organization{ID: defaultOrgID}, nil + } + cache := newChatConfigCache(ctx, store, clock) + + // A miss inside the default org must not recurse into the fallback: + // the default-org resolution happens once, the self-check short + // circuits, and the original miss propagates. + _, err := cache.DefaultModelConfig(ctx, defaultOrgID) + require.ErrorIs(t, err, sql.ErrNoRows) + require.Equal(t, int32(1), store.defaultOrganizationCall.Load()) + require.Equal(t, int32(1), store.defaultModelConfigCallCount(defaultOrgID)) +} + +func TestConfigCache_DefaultModelConfig_FallbackMiss(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + clock := quartz.NewMock(t) + defaultOrgID := uuid.New() + otherOrgID := uuid.New() store := &stubChatConfigStore{} store.getDefaultChatModelConfig = func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { return database.ChatModelConfig{}, sql.ErrNoRows } + store.getDefaultOrganization = func(context.Context) (database.Organization, error) { + return database.Organization{ID: defaultOrgID}, nil + } cache := newChatConfigCache(ctx, store, clock) - // A miss inside the org propagates unchanged. - _, err := cache.DefaultModelConfig(ctx, orgID) + // Neither the org nor the default org has a default: the miss + // propagates unchanged. + _, err := cache.DefaultModelConfig(ctx, otherOrgID) require.ErrorIs(t, err, sql.ErrNoRows) - require.Equal(t, int32(1), store.defaultModelConfigCallCount(orgID)) + require.Equal(t, int32(1), store.defaultModelConfigCallCount(otherOrgID)) + require.Equal(t, int32(1), store.defaultModelConfigCallCount(defaultOrgID)) } func TestConfigCache_UserPrompt_NegativeCaching(t *testing.T) { diff --git a/coderd/x/chatd/subagent.go b/coderd/x/chatd/subagent.go index e115eb56b34..04a560a48f7 100644 --- a/coderd/x/chatd/subagent.go +++ b/coderd/x/chatd/subagent.go @@ -630,7 +630,7 @@ func (p *Server) listSpawnableModelConfigs( ) ([]map[string]any, error) { //nolint:gocritic // Chatd needs its scoped config and user-data access here. chatdCtx := dbauthz.AsChatd(ctx) - rows, err := p.db.GetEnabledChatModelConfigsByOrganization(chatdCtx, organizationID) + rows, err := enabledChatModelConfigsWithDefaultOrgFallback(chatdCtx, p.db, organizationID) if err != nil { return nil, xerrors.Errorf("get enabled chat model configs: %w", err) } @@ -1265,16 +1265,6 @@ func (p *Server) createChildSubagentChatWithOptions( if opts.modelConfigIDOverride != nil { modelConfigID = *opts.modelConfigIDOverride } - if modelConfigID != uuid.Nil && modelConfigID != parent.LastModelConfigID { - modelConfig, err := p.configCache.ModelConfigByID(ctx, modelConfigID) - if err != nil { - return database.Chat{}, xerrors.Errorf("get child model config: %w", err) - } - if modelConfig.OrganizationID != parent.OrganizationID { - modelConfigID = parent.LastModelConfigID - opts.reasoningEffortOverride = nil - } - } if modelConfigID == uuid.Nil { return database.Chat{}, xerrors.New("model config is required") } diff --git a/coderd/x/chatd/subagent_internal_test.go b/coderd/x/chatd/subagent_internal_test.go index 512f901716c..5ecd505ff94 100644 --- a/coderd/x/chatd/subagent_internal_test.go +++ b/coderd/x/chatd/subagent_internal_test.go @@ -1554,40 +1554,6 @@ func TestCreateChildSubagentChat_OverrideWorksWhenParentHasNoModel(t *testing.T) require.Equal(t, overrideModel.ID, childChat.LastModelConfigID) } -func TestCreateChildSubagentChat_CrossOrgOverrideFallsBackToParent(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{}) - - ctx := chatdTestContext(t) - user, org, parentModel := seedInternalChatDeps(t, db) - otherOrg := dbgen.Organization(t, db, database.Organization{}) - overrideModel := insertInternalChatModelConfig( - t, db, otherOrg.ID, "cross-org-child-"+uuid.NewString(), true, - ) - parentChat := createInternalParentChat( - ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-cross-org-override", - ) - - child, err := server.createChildSubagentChatWithOptions( - ctx, - parentChat, - "delegate work", - "", - childSubagentChatOptions{ - modelConfigIDOverride: &overrideModel.ID, - reasoningEffortOverride: ptr.Ref("high"), - }, - ) - require.NoError(t, err) - - childChat, err := db.GetChatByID(ctx, child.ID) - require.NoError(t, err) - require.Equal(t, parentModel.ID, childChat.LastModelConfigID) - require.False(t, childChat.LastReasoningEffort.Valid) -} - func TestSpawnAgent_ExplicitModelConfigID(t *testing.T) { t.Parallel() @@ -4893,7 +4859,205 @@ func TestListAgents(t *testing.T) { }) } -func TestListSubagentModels_NonDefaultOrgSeesOnlyOwnOrgConfigs(t *testing.T) { +type enabledChatModelConfigsReadSkewStore struct { + database.Store + organizationID uuid.UUID + config database.ChatModelConfig + listCalls atomic.Int32 +} + +func (s *enabledChatModelConfigsReadSkewStore) GetEnabledChatModelConfigsByOrganization( + _ context.Context, + organizationID uuid.UUID, +) ([]database.GetEnabledChatModelConfigsByOrganizationRow, error) { + if organizationID != s.organizationID { + return nil, sql.ErrNoRows + } + if s.listCalls.Add(1) == 1 { + return nil, nil + } + return []database.GetEnabledChatModelConfigsByOrganizationRow{{ChatModelConfig: s.config}}, nil +} + +func (s *enabledChatModelConfigsReadSkewStore) GetDefaultChatModelConfig( + _ context.Context, + organizationID uuid.UUID, +) (database.ChatModelConfig, error) { + if organizationID != s.organizationID { + return database.ChatModelConfig{}, sql.ErrNoRows + } + return s.config, nil +} + +type enabledChatModelConfigsEmptyDefaultOrgStore struct { + database.Store + organizationID uuid.UUID +} + +func (s *enabledChatModelConfigsEmptyDefaultOrgStore) GetEnabledChatModelConfigsByOrganization( + _ context.Context, + organizationID uuid.UUID, +) ([]database.GetEnabledChatModelConfigsByOrganizationRow, error) { + if organizationID != s.organizationID { + return nil, xerrors.Errorf("unexpected organization %q", organizationID) + } + return nil, nil +} + +func (s *enabledChatModelConfigsEmptyDefaultOrgStore) GetDefaultChatModelConfig( + _ context.Context, + organizationID uuid.UUID, +) (database.ChatModelConfig, error) { + if organizationID != s.organizationID { + return database.ChatModelConfig{}, xerrors.Errorf("unexpected organization %q", organizationID) + } + return database.ChatModelConfig{}, sql.ErrNoRows +} + +func (s *enabledChatModelConfigsEmptyDefaultOrgStore) GetDefaultOrganization(context.Context) (database.Organization, error) { + return database.Organization{ID: s.organizationID}, nil +} + +func TestEnabledChatModelConfigsWithDefaultOrgFallback(t *testing.T) { + t.Parallel() + + db, _ := dbtestutil.NewDB(t) + ctx := testutil.Context(t, testutil.WaitShort) + + defaultOrg, err := db.GetDefaultOrganization(ctx) + require.NoError(t, err) + otherOrg := dbgen.Organization(t, db, database.Organization{}) + provider := dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{ + Type: database.AIProviderTypeOpenai, + }, "test-key") + // Every write path promotes a default within the org it writes to, so + // an org that owns configs always owns one. Fixtures that insert + // directly must uphold that: the fallback reads it as the ownership + // marker. + defaultOrgConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, + OrganizationID: defaultOrg.ID, + IsDefault: true, + }) + + t.Run("OrgListEmptyFallsBackToDefaultOrg", func(t *testing.T) { + t.Parallel() + + // A fresh org is guaranteed empty even when the test database + // carries seeded configs in previously created orgs. + emptyOrg := dbgen.Organization(t, db, database.Organization{}) + ctx := testutil.Context(t, testutil.WaitShort) + rows, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, db, emptyOrg.ID) + require.NoError(t, err) + require.NotEmpty(t, rows) + + found := slices.ContainsFunc(rows, func(row database.GetEnabledChatModelConfigsByOrganizationRow) bool { + return row.ChatModelConfig.ID == defaultOrgConfig.ID + }) + require.True(t, found, "default org list should include its config") + }) + + t.Run("ReFetchesWhenDefaultAppearsAfterEmptyList", func(t *testing.T) { + t.Parallel() + + organizationID := uuid.New() + config := database.ChatModelConfig{ + ID: uuid.New(), + OrganizationID: organizationID, + IsDefault: true, + } + store := &enabledChatModelConfigsReadSkewStore{ + organizationID: organizationID, + config: config, + } + + rows, err := enabledChatModelConfigsWithDefaultOrgFallback(t.Context(), store, organizationID) + require.NoError(t, err) + require.Len(t, rows, 1) + require.Equal(t, config.ID, rows[0].ChatModelConfig.ID) + require.EqualValues(t, 2, store.listCalls.Load()) + }) + + t.Run("OrgListPresentNeverFallsBack", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + + // Give the other org its own enabled config: the default org's + // list must not leak in. + ownConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, + OrganizationID: otherOrg.ID, + }) + + rows, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, db, otherOrg.ID) + require.NoError(t, err) + require.Len(t, rows, 1) + require.Equal(t, ownConfig.ID, rows[0].ChatModelConfig.ID) + }) + + t.Run("DefaultOrgNeverFallsBack", func(t *testing.T) { + t.Parallel() + + store := &enabledChatModelConfigsEmptyDefaultOrgStore{ + Store: db, + organizationID: defaultOrg.ID, + } + got, err := enabledChatModelConfigsWithDefaultOrgFallback(t.Context(), store, defaultOrg.ID) + require.NoError(t, err) + require.Empty(t, got) + }) + + t.Run("OrgWithEveryConfigDisabledNeverFallsBack", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + + // The org owns a config but has disabled it, so its enabled + // list is empty. It must keep that empty list instead of + // borrowing the default org's models. + disabledOrg := dbgen.Organization(t, db, database.Organization{}) + dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, + OrganizationID: disabledOrg.ID, + IsDefault: true, + }, func(params *database.InsertChatModelConfigParams) { + // dbgen defaults Enabled to true, so disable it here. + params.Enabled = false + }) + + got, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, db, disabledOrg.ID) + require.NoError(t, err) + require.Empty(t, got) + }) + + t.Run("OrgWithDisabledProviderNeverFallsBack", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + + // The org's only config is enabled but its provider is + // disabled, which also yields an empty enabled list. + disabledProviderOrg := dbgen.Organization(t, db, database.Organization{}) + disabledProvider := dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{ + Type: database.AIProviderTypeOpenai, + }, "test-key", func(params *database.InsertAIProviderParams) { + // dbgen defaults Enabled to true, so disable it here. + params.Enabled = false + }) + dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + AIProviderID: uuid.NullUUID{UUID: disabledProvider.ID, Valid: true}, + OrganizationID: disabledProviderOrg.ID, + IsDefault: true, + }) + + got, err := enabledChatModelConfigsWithDefaultOrgFallback(ctx, db, disabledProviderOrg.ID) + require.NoError(t, err) + require.Empty(t, got) + }) +} + +func TestListSubagentModels_NonDefaultOrgListIsOrgLocal(t *testing.T) { t.Parallel() db, ps := dbtestutil.NewDB(t) @@ -4904,7 +5068,8 @@ func TestListSubagentModels_NonDefaultOrgSeesOnlyOwnOrgConfigs(t *testing.T) { // The chat's org has its own config (the seed model), so the // list is org-local and the default org's config must not leak - // in. + // in. The empty-org fallback is covered by + // TestEnabledChatModelConfigsWithDefaultOrgFallback. defaultOrg, err := db.GetDefaultOrganization(ctx) require.NoError(t, err) defaultOrgProvider := dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{ diff --git a/coderd/x/chatd/title_override.go b/coderd/x/chatd/title_override.go index bb2a0ddba46..4056fdcfe13 100644 --- a/coderd/x/chatd/title_override.go +++ b/coderd/x/chatd/title_override.go @@ -83,7 +83,7 @@ func (p *Server) resolveTitleGenerationModelOverride( if err != nil { return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, overrideSet, err } - if !overrideSet || modelConfig.OrganizationID != chat.OrganizationID { + if !overrideSet { return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, false, nil } modelConfig = withResolvedReasoningEffort(modelConfig, overrideEffort) diff --git a/coderd/x/chatd/title_override_internal_test.go b/coderd/x/chatd/title_override_internal_test.go index e52a1b7240d..7df190a4c40 100644 --- a/coderd/x/chatd/title_override_internal_test.go +++ b/coderd/x/chatd/title_override_internal_test.go @@ -401,7 +401,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnset(t *testing.T) { require.Equal(t, preferredConfig, gotConfig) } -func TestResolveManualTitleModel_CrossOrgConfigsAreInvisible(t *testing.T) { +func TestResolveManualTitleModel_NonDefaultOrgUsesDefaultOrgConfigs(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitShort) @@ -409,11 +409,29 @@ func TestResolveManualTitleModel_CrossOrgConfigsAreInvisible(t *testing.T) { db := dbmock.NewMockStore(ctrl) logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) chat, _ := titleOverrideTestChatAndMessages(t) - chat.OrganizationID = uuid.New() + chat.OrganizationID = uuid.New() // non-default org, no configs of its own + defaultOrgID := uuid.New() + providerID := uuid.New() + preferredConfig := database.ChatModelConfig{ + ID: uuid.New(), + AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, + Model: preferredTitleModels[1].model, + Enabled: true, + OrganizationID: defaultOrgID, + } db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil) + // The chat's org has no enabled configs; the selector must fall + // back to the default org's list until the org-scoping cutover. db.EXPECT().GetEnabledChatModelConfigsByOrganization(gomock.Any(), chat.OrganizationID).Return(nil, nil) - db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), chat.OrganizationID).Return(database.ChatModelConfig{}, sql.ErrNoRows) + db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), chat.OrganizationID). + Return(database.ChatModelConfig{}, sql.ErrNoRows) + db.EXPECT().GetDefaultOrganization(gomock.Any()).Return(database.Organization{ID: defaultOrgID}, nil) + db.EXPECT().GetDefaultChatModelConfig(gomock.Any(), defaultOrgID).Return(preferredConfig, nil) + db.EXPECT().GetEnabledChatModelConfigsByOrganization(gomock.Any(), defaultOrgID).Return([]database.GetEnabledChatModelConfigsByOrganizationRow{ + {ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider}, + }, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() server := titleOverrideTestServer(db, logger) model, gotConfig, err := server.resolveManualTitleModel( @@ -422,9 +440,9 @@ func TestResolveManualTitleModel_CrossOrgConfigsAreInvisible(t *testing.T) { chat, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) - require.ErrorIs(t, err, ErrNoDefaultChatModelConfig) - require.False(t, model.Valid()) - require.Equal(t, database.ChatModelConfig{}, gotConfig) + require.NoError(t, err) + require.NotNil(t, model) + require.Equal(t, preferredConfig, gotConfig) } func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testing.T) { @@ -543,37 +561,6 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUsable(t *testing.T) require.Equal(t, overrideConfig, gotConfig) } -func TestResolveTitleGenerationModelOverride_CrossOrgFallsBack(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - ctrl := gomock.NewController(t) - db := dbmock.NewMockStore(ctrl) - logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) - chat, _ := titleOverrideTestChatAndMessages(t) - chat.OrganizationID = uuid.New() - overrideConfig := titleOverrideModelConfig("gpt-4.1", true) - overrideConfig.OrganizationID = uuid.New() - providerID := uuid.New() - overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} - - db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) - db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) - db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() - db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ - ProviderID: providerID, - APIKey: "test-key", - }}, nil).AnyTimes() - - server := titleOverrideTestServer(db, logger) - modelConfig, model, route, overrideSet, err := server.resolveTitleGenerationModelOverride(ctx, chat, modelBuildOptions{}) - require.NoError(t, err) - require.False(t, overrideSet) - require.Equal(t, database.ChatModelConfig{}, modelConfig) - require.False(t, model.Valid()) - require.Equal(t, aiGatewayModelRoute{}, route) -} - func TestResolveManualTitleModel_TitleGenerationOverrideMissingCredentials(t *testing.T) { t.Parallel() diff --git a/enterprise/coderd/exp_chats_test.go b/enterprise/coderd/exp_chats_test.go index 3f9b9196963..364b82fc682 100644 --- a/enterprise/coderd/exp_chats_test.go +++ b/enterprise/coderd/exp_chats_test.go @@ -14,7 +14,6 @@ import ( "github.com/coder/coder/v2/coderd/aibridgedtest" "github.com/coder/coder/v2/coderd/coderdtest" "github.com/coder/coder/v2/coderd/database" - "github.com/coder/coder/v2/coderd/database/dbgen" "github.com/coder/coder/v2/coderd/database/dbtestutil" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/util/ptr" @@ -1080,7 +1079,7 @@ func TestCreateChatNonDefaultOrg(t *testing.T) { ctx := testutil.Context(t, testutil.WaitLong) - client, db, firstUser := coderdenttest.NewWithDatabase(t, &coderdenttest.Options{ + client, firstUser := coderdenttest.New(t, &coderdenttest.Options{ Options: &coderdtest.Options{ DeploymentValues: func() *codersdk.DeploymentValues { v := coderdtest.DeploymentValues(t) @@ -1108,13 +1107,6 @@ func TestCreateChatNonDefaultOrg(t *testing.T) { // Create a second (non-default) org via the API. secondOrg := coderdenttest.CreateOrganization(t, client, coderdenttest.CreateOrganizationOptions{}) - dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, - Model: "gpt-4o-mini", - Enabled: true, - IsDefault: true, - OrganizationID: secondOrg.ID, - }) // Create a member with agents-access in both orgs. memberClientRaw, member := coderdtest.CreateAnotherUser( @@ -1156,7 +1148,7 @@ func TestListChats_OrgAdminOnlySeesOwnChats(t *testing.T) { ctx := testutil.Context(t, testutil.WaitLong) - client, db, firstUser := coderdenttest.NewWithDatabase(t, &coderdenttest.Options{ + client, firstUser := coderdenttest.New(t, &coderdenttest.Options{ Options: &coderdtest.Options{ DeploymentValues: func() *codersdk.DeploymentValues { v := coderdtest.DeploymentValues(t) @@ -1184,13 +1176,6 @@ func TestListChats_OrgAdminOnlySeesOwnChats(t *testing.T) { // Create a second (non-default) org. secondOrg := coderdenttest.CreateOrganization(t, client, coderdenttest.CreateOrganizationOptions{}) - dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, - Model: "gpt-4o-mini", - Enabled: true, - IsDefault: true, - OrganizationID: secondOrg.ID, - }) // Create a member with agents-access in both orgs. memberClientRaw, _ := coderdtest.CreateAnotherUser(