diff --git a/apps/web/src/app/admin/custom-llms/CustomLlmsContent.tsx b/apps/web/src/app/admin/custom-llms/CustomLlmsContent.tsx index f35f0fb2a1..d43c3b6d95 100644 --- a/apps/web/src/app/admin/custom-llms/CustomLlmsContent.tsx +++ b/apps/web/src/app/admin/custom-llms/CustomLlmsContent.tsx @@ -53,6 +53,7 @@ const INITIAL_DEFINITION: CustomLlmDefinition = { max_completion_tokens: 0, base_url: '', organization_ids: [], + group_ids: [], }; const INITIAL_CREDENTIALS: CustomLlmCredentials = { diff --git a/apps/web/src/components/organizations/groups/policies/model-access/ModelAccessPolicyEditor.tsx b/apps/web/src/components/organizations/groups/policies/model-access/ModelAccessPolicyEditor.tsx index 00d702e3d7..0c24ecd3fb 100644 --- a/apps/web/src/components/organizations/groups/policies/model-access/ModelAccessPolicyEditor.tsx +++ b/apps/web/src/components/organizations/groups/policies/model-access/ModelAccessPolicyEditor.tsx @@ -307,7 +307,7 @@ export function ModelAccessPolicyEditor({ )}

- Direct BYOK models and custom LLMs remain available organization-wide. + Direct BYOK models remain available organization-wide.

{ + it('allows organization-wide access without a matching group', () => { + expect(hasCustomLlmAccess(definition, 'organization-1', [])).toBe(true); + }); + + it('allows access through a matching group when the organization is not allowed', () => { + expect( + hasCustomLlmAccess(definition, 'organization-2', ['00000000-0000-4000-8000-000000000001']) + ).toBe(true); + }); + + it('denies access when neither the organization nor a group matches', () => { + expect( + hasCustomLlmAccess(definition, 'organization-2', ['00000000-0000-4000-8000-000000000002']) + ).toBe(false); + }); +}); diff --git a/apps/web/src/lib/ai-gateway/custom-llm/access.ts b/apps/web/src/lib/ai-gateway/custom-llm/access.ts new file mode 100644 index 0000000000..5934805034 --- /dev/null +++ b/apps/web/src/lib/ai-gateway/custom-llm/access.ts @@ -0,0 +1,42 @@ +import { organization_group_memberships } from '@kilocode/db/schema'; +import type { CustomLlmDefinition } from '@kilocode/db/schema-types'; +import { readDb } from '@/lib/drizzle'; +import { and, eq, inArray } from 'drizzle-orm'; + +export function hasCustomLlmAccess( + definition: CustomLlmDefinition, + organizationId: string, + groupIds: readonly string[] +) { + return ( + definition.organization_ids.includes(organizationId) || + definition.group_ids?.some(groupId => groupIds.includes(groupId)) === true + ); +} + +export async function userHasCustomLlmAccess( + definition: CustomLlmDefinition, + organizationId: string, + kiloUserId: string +) { + if (hasCustomLlmAccess(definition, organizationId, [])) { + return true; + } + if (!definition.group_ids?.length) { + return false; + } + + const [membership] = await readDb + .select({ groupId: organization_group_memberships.group_id }) + .from(organization_group_memberships) + .where( + and( + eq(organization_group_memberships.organization_id, organizationId), + eq(organization_group_memberships.kilo_user_id, kiloUserId), + inArray(organization_group_memberships.group_id, definition.group_ids) + ) + ) + .limit(1); + + return Boolean(membership); +} diff --git a/apps/web/src/lib/ai-gateway/custom-llm/listAvailableCustomLlms.ts b/apps/web/src/lib/ai-gateway/custom-llm/listAvailableCustomLlms.ts index 2e67b4b29e..16907ddc96 100644 --- a/apps/web/src/lib/ai-gateway/custom-llm/listAvailableCustomLlms.ts +++ b/apps/web/src/lib/ai-gateway/custom-llm/listAvailableCustomLlms.ts @@ -2,6 +2,7 @@ import { custom_llm2 } from '@kilocode/db/schema'; import { readDb } from '@/lib/drizzle'; import { CustomLlmDefinitionSchema, type CustomLlmDefinition } from '@kilocode/db/schema-types'; import { orderOpenCodeSettings } from './order-opencode-variants'; +import { hasCustomLlmAccess } from './access'; function convert(publicId: string, model: CustomLlmDefinition) { return { @@ -41,7 +42,7 @@ function convert(publicId: string, model: CustomLlmDefinition) { }; } -export async function listAvailableCustomLlms(organizationId: string) { +export async function listAvailableCustomLlms(organizationId: string, groupIds: readonly string[]) { const rows = await readDb.select().from(custom_llm2); return rows .map(row => { @@ -52,6 +53,6 @@ export async function listAvailableCustomLlms(organizationId: string) { return parsed.success ? { public_id: row.public_id, definition: parsed.data } : null; }) .filter(row => row !== null) - .filter(row => row.definition.organization_ids.includes(organizationId)) + .filter(row => hasCustomLlmAccess(row.definition, organizationId, groupIds)) .map(row => convert(row.public_id, row.definition)); } diff --git a/apps/web/src/lib/ai-gateway/providers/get-provider.ts b/apps/web/src/lib/ai-gateway/providers/get-provider.ts index 86dce64c4a..4d5790417d 100644 --- a/apps/web/src/lib/ai-gateway/providers/get-provider.ts +++ b/apps/web/src/lib/ai-gateway/providers/get-provider.ts @@ -23,6 +23,7 @@ import { type AllocationSubject, } from '@/lib/ai-gateway/experiments/pick-variant'; import { getGoogleServiceAccountAccessToken } from '@/lib/ai-gateway/custom-llm/google-service-account'; +import { userHasCustomLlmAccess } from '@/lib/ai-gateway/custom-llm/access'; import { decryptApiKey } from '@/lib/ai-gateway/byok/encryption'; import { BYOK_ENCRYPTION_KEY } from '@/lib/config.server'; @@ -95,7 +96,8 @@ async function checkDirectBYOK( async function checkCustomLlm( requestedModel: string, - organizationId: string + organizationId: string, + kiloUserId: string ): Promise { const [row] = await readDb .select() @@ -106,7 +108,7 @@ async function checkCustomLlm( console.log('Failed to parse custom llm definition', parsedCustomLlm.error); } const customLlm = parsedCustomLlm.data; - if (!customLlm || !customLlm.organization_ids.includes(organizationId)) { + if (!customLlm || !(await userHasCustomLlmAccess(customLlm, organizationId, kiloUserId))) { return null; } @@ -251,8 +253,8 @@ export async function getProvider(input: GetProviderInput): Promise group.id); groupPolicies = groups.map( group => parsePolicies(group.policies, { @@ -180,6 +183,7 @@ export async function getOrganizationGroupPolicyContext(params: { return { organization, defaultPolicies, + groupIds, groupPolicies, policyRevision: settings?.policy_revision ?? 0, }; diff --git a/apps/web/src/lib/organizations/organization-models.ts b/apps/web/src/lib/organizations/organization-models.ts index b2e531c72a..a5242445d9 100644 --- a/apps/web/src/lib/organizations/organization-models.ts +++ b/apps/web/src/lib/organizations/organization-models.ts @@ -55,7 +55,7 @@ export async function getAvailableModelsForOrganization( } availableModels.push(...(await getDirectByokModelsForOrganization(organizationId))); - availableModels.push(...(await listAvailableCustomLlms(organizationId))); + availableModels.push(...(await listAvailableCustomLlms(organizationId, context.groupIds))); return { ...responseData, diff --git a/apps/web/src/routers/admin/custom-llm-router.test.ts b/apps/web/src/routers/admin/custom-llm-router.test.ts index d4fc7df952..2665725d80 100644 --- a/apps/web/src/routers/admin/custom-llm-router.test.ts +++ b/apps/web/src/routers/admin/custom-llm-router.test.ts @@ -18,6 +18,7 @@ const validDefinition: CustomLlmDefinition = { max_completion_tokens: 4096, base_url: 'https://api.openai.com/v1', organization_ids: ['org_test_123'], + group_ids: ['00000000-0000-4000-8000-000000000123'], }; beforeEach(async () => { @@ -52,6 +53,7 @@ describe('adminCustomLlmRouter', () => { expect(result.public_id).toBe(publicId); expect(result.definition.display_name).toBe('Custom GPT-4'); + expect(result.definition.group_ids).toEqual(validDefinition.group_ids); expect((result.definition as Record).api_key).toBeUndefined(); expect((result as Record).encrypted_api_key).toBeUndefined(); diff --git a/apps/web/src/routers/organizations/organization-settings-router.test.ts b/apps/web/src/routers/organizations/organization-settings-router.test.ts index b2490caad3..2470e5401c 100644 --- a/apps/web/src/routers/organizations/organization-settings-router.test.ts +++ b/apps/web/src/routers/organizations/organization-settings-router.test.ts @@ -13,7 +13,10 @@ import type { import { type User, type Organization, + custom_llm2, organization_audit_logs, + organization_group_memberships, + organization_groups, organizations, } from '@kilocode/db/schema'; import { eq } from 'drizzle-orm'; @@ -54,6 +57,7 @@ import { getEnhancedOpenRouterModels } from '@/lib/ai-gateway/providers/openrout import { getProviderSlugsForModel } from '@/lib/ai-gateway/providers/openrouter/models-by-provider-index.server'; import { isPublicIdExperimented } from '@/lib/ai-gateway/experiments/membership'; import { CLAUDE_SONNET_LATEST_MODEL_ALIAS } from '@/lib/ai-gateway/latest-model-aliases'; +import { userHasCustomLlmAccess } from '@/lib/ai-gateway/custom-llm/access'; function makeTestOpenRouterModel(id: string): OpenRouterModel { return { @@ -351,6 +355,68 @@ describe('organizations settings trpc router', () => { }; } + it('includes group-only custom LLMs only for members of an allowed group', async () => { + const organization = await createTestOrganization( + 'Custom LLM Group Access', + owner.id, + 0, + {}, + false + ); + await addUserToOrganization(organization.id, member.id, 'member'); + const [group] = await db + .insert(organization_groups) + .values({ + organization_id: organization.id, + name: `Custom LLM ${randomUUID()}`, + }) + .returning(); + await db.insert(organization_group_memberships).values({ + organization_id: organization.id, + group_id: group.id, + kilo_user_id: member.id, + }); + + const publicId = `kilo-internal/group-only-${randomUUID()}`; + const definition = { + internal_id: 'group-only-upstream', + display_name: 'Group-only custom LLM', + context_length: 128_000, + max_completion_tokens: 4096, + base_url: 'https://example.com/v1', + organization_ids: [], + group_ids: [group.id], + }; + await db.insert(custom_llm2).values({ + public_id: publicId, + definition, + }); + + try { + await expect(userHasCustomLlmAccess(definition, organization.id, member.id)).resolves.toBe( + true + ); + await expect(userHasCustomLlmAccess(definition, organization.id, owner.id)).resolves.toBe( + false + ); + + const memberCaller = await createCallerForUser(member.id); + const memberResult = await memberCaller.organizations.settings.listAvailableModels({ + organizationId: organization.id, + }); + const ownerCaller = await createCallerForUser(owner.id); + const ownerResult = await ownerCaller.organizations.settings.listAvailableModels({ + organizationId: organization.id, + }); + + expect(memberResult.data.some(model => model.id === publicId)).toBe(true); + expect(ownerResult.data.some(model => model.id === publicId)).toBe(false); + } finally { + await db.delete(custom_llm2).where(eq(custom_llm2.public_id, publicId)); + await db.delete(organizations).where(eq(organizations.id, organization.id)); + } + }); + it('excludes models outside the snapshot without configured restrictions', async () => { const organization = await createTestOrganization( 'Snapshot-only Enterprise', diff --git a/packages/db/src/schema-types.ts b/packages/db/src/schema-types.ts index 605492f0c7..faed671ce6 100644 --- a/packages/db/src/schema-types.ts +++ b/packages/db/src/schema-types.ts @@ -1978,6 +1978,7 @@ export const CustomLlmDefinitionSchema = z.object({ ...CustomLlmApiConfigSchema.shape, display_name: z.string(), organization_ids: z.array(z.string()), + group_ids: z.array(z.uuid()).optional(), pricing: CustomLlmPricingSchema.optional(), });