Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
81 changes: 81 additions & 0 deletions packages/types/src/__tests__/organization-allow-list.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
import { describe, expect, it } from "vitest"

import type { OrganizationAllowList } from "../cloud.js"
import { isModelAllowedForOrganization, isProviderAllowAll } from "../organization-allow-list.js"
import { providerIdentifiers } from "../provider-identifiers.js"

describe("isProviderAllowAll", () => {
it("allows everything when the allow list is missing", () => {
expect(isProviderAllowAll(undefined, providerIdentifiers.anthropic)).toBe(true)
})

it("allows everything when the organization allows all", () => {
const allowList: OrganizationAllowList = { allowAll: true, providers: {} }
expect(isProviderAllowAll(allowList, providerIdentifiers.anthropic)).toBe(true)
})

it("rejects when the provider is missing but the organization is restricted", () => {
const allowList: OrganizationAllowList = { allowAll: false, providers: {} }
expect(isProviderAllowAll(allowList, undefined)).toBe(false)
})

it("rejects a provider without an explicit entry", () => {
const allowList: OrganizationAllowList = { allowAll: false, providers: {} }
expect(isProviderAllowAll(allowList, providerIdentifiers.anthropic)).toBe(false)
})

it("allows a provider narrowed to a model list (the list still constrains models)", () => {
const allowList: OrganizationAllowList = {
allowAll: false,
providers: { [providerIdentifiers.anthropic]: { allowAll: false, models: ["m"] } },
}
// A non-allowAll provider must not expose the custom-model escape hatch.
expect(isProviderAllowAll(allowList, providerIdentifiers.anthropic)).toBe(false)
})

it("allows a provider explicitly marked allowAll", () => {
const allowList: OrganizationAllowList = {
allowAll: false,
providers: { [providerIdentifiers.openrouter]: { allowAll: true } },
}
expect(isProviderAllowAll(allowList, providerIdentifiers.openrouter)).toBe(true)
})
})

describe("isModelAllowedForOrganization", () => {
it("allows any model when the allow list is missing", () => {
expect(isModelAllowedForOrganization(undefined, providerIdentifiers.anthropic, "any")).toBe(true)
})

it("allows any model when the organization allows all", () => {
const allowList: OrganizationAllowList = { allowAll: true, providers: {} }
expect(isModelAllowedForOrganization(allowList, providerIdentifiers.anthropic, "any")).toBe(true)
})

it("rejects when the provider is missing but the organization is restricted", () => {
const allowList: OrganizationAllowList = { allowAll: false, providers: {} }
expect(isModelAllowedForOrganization(allowList, undefined, "any")).toBe(false)
})

it("rejects a model from a provider without an explicit entry", () => {
const allowList: OrganizationAllowList = { allowAll: false, providers: {} }
expect(isModelAllowedForOrganization(allowList, providerIdentifiers.anthropic, "m")).toBe(false)
})

it("allows a listed model and rejects an unlisted one", () => {
const allowList: OrganizationAllowList = {
allowAll: false,
providers: { [providerIdentifiers.anthropic]: { allowAll: false, models: ["allowed"] } },
}
expect(isModelAllowedForOrganization(allowList, providerIdentifiers.anthropic, "allowed")).toBe(true)
expect(isModelAllowedForOrganization(allowList, providerIdentifiers.anthropic, "other")).toBe(false)
})

it("allows any model for a provider explicitly marked allowAll", () => {
const allowList: OrganizationAllowList = {
allowAll: false,
providers: { [providerIdentifiers.openrouter]: { allowAll: true } },
}
expect(isModelAllowedForOrganization(allowList, providerIdentifiers.openrouter, "anything")).toBe(true)
})
})
1 change: 1 addition & 0 deletions packages/types/src/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ export * from "./mcp.js"
export * from "./message.js"
export * from "./mode.js"
export * from "./model.js"
export * from "./organization-allow-list.js"
export * from "./provider-identifiers.js"
export * from "./provider-settings.js"
export * from "./task.js"
Expand Down
54 changes: 54 additions & 0 deletions packages/types/src/organization-allow-list.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
import type { OrganizationAllowList } from "./cloud.js"
import type { ProviderName } from "./provider-settings.js"

/**
* Whether organization policy allows an arbitrary/custom model id for a provider.
*
* Only providers that are explicitly `allowAll` may use free-form ids; when a
* provider is narrowed to an allow-list, callers must not expose a custom-model
* escape hatch. This is the single source of truth shared by the settings
* model picker and the chat model selector.
*/
export const isProviderAllowAll = (
allowList: OrganizationAllowList | undefined,
provider: ProviderName | undefined,
): boolean => {
if (!allowList || allowList.allowAll) {
return true
}

if (!provider) {

Check warning on line 20 in packages/types/src/organization-allow-list.ts

View workflow job for this annotation

GitHub Actions / mutation-diff

Mutation test advisory

packages/types/src/organization-allow-list.ts:20: 2 mutation test gaps; example: Survived ConditionalExpression mutant (replacement: false). See the job summary for the complete list and resolution guidance.
return false
}

return allowList.providers[provider]?.allowAll === true
}

/**
* Whether organization policy allows selecting the given model id for a provider.
*
* `allowAll` (org-wide or per provider) permits any id; otherwise the id must be
* present in the provider's explicit `models` list. Intended as a defense-in-depth
* check that rejects a selection regardless of how it was produced (option click,
* keyboard navigation, or custom entry).
*/
export const isModelAllowedForOrganization = (
allowList: OrganizationAllowList | undefined,
provider: ProviderName | undefined,
modelId: string,
): boolean => {
if (!allowList || allowList.allowAll) {
return true
}

if (!provider) {

Check warning on line 44 in packages/types/src/organization-allow-list.ts

View workflow job for this annotation

GitHub Actions / mutation-diff

Mutation test advisory

packages/types/src/organization-allow-list.ts:44: 2 mutation test gaps; example: Survived ConditionalExpression mutant (replacement: false). See the job summary for the complete list and resolution guidance.
return false
}

const providerConfig = allowList.providers[provider]
if (!providerConfig) {
return false
}

return providerConfig.allowAll === true || (providerConfig.models?.includes(modelId) ?? false)

Check warning on line 53 in packages/types/src/organization-allow-list.ts

View workflow job for this annotation

GitHub Actions / mutation-diff

Mutation test advisory

packages/types/src/organization-allow-list.ts:53: 2 mutation test gaps; example: NoCoverage BooleanLiteral mutant (replacement: true). See the job summary for the complete list and resolution guidance.
}
1 change: 1 addition & 0 deletions packages/types/src/vscode-extension-host.ts
Original file line number Diff line number Diff line change
Expand Up @@ -464,6 +464,7 @@ export type EditQueuedMessagePayload = Pick<QueuedMessage, "id" | "text" | "imag

export interface WebviewMessage {
type:
| "cancelModelRequest"
| "updateTodoList"
| "deleteMultipleTasksWithIds"
| "currentApiConfigName"
Expand Down
17 changes: 17 additions & 0 deletions src/api/providers/__tests__/openai.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2008,6 +2008,23 @@ describe("getOpenAiModels", () => {
vi.mocked(axios.get).mockClear()
})

it("forwards cancellation to the HTTP request", async () => {
const controller = new AbortController()
vi.mocked(axios.get).mockImplementationOnce(
(_url, options) =>
new Promise((_resolve, reject) => {
options?.signal?.addEventListener?.("abort", () => reject(new Error("cancelled")))
}),
)
const pending = getOpenAiModels("https://example.test/v1", "test-key", undefined, controller.signal)
expect(axios.get).toHaveBeenCalledWith(
"https://example.test/v1/models",
expect.objectContaining({ signal: controller.signal }),
)
controller.abort()
expect(await pending).toEqual([])
})

it("should return empty array when baseUrl is not provided", async () => {
const result = await getOpenAiModels(undefined, "test-key")
expect(result).toEqual([])
Expand Down
23 changes: 22 additions & 1 deletion src/api/providers/fetchers/__tests__/kimi-code.spec.ts
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import { getEventListeners } from "events"

import { getKimiCodeModels, kimiCodeModelSchema, mapKimiCodeModel } from "../kimi-code"

describe("Kimi Code model discovery", () => {
Expand Down Expand Up @@ -83,18 +85,37 @@ describe("Kimi Code model discovery", () => {
expect(models["model-b"].supportsReasoningEffort).toEqual(["low", "high", "max"])
})

it.each(["success", "failure"])("cleans up caller listeners and timeout on %s", async (outcome) => {
vi.useFakeTimers()
const controller = new AbortController()
const transport = vi.spyOn(globalThis, "fetch")
const failure = new Error("network failed")
if (outcome === "success") transport.mockResolvedValue(new Response(JSON.stringify({ data: [] })))
else transport.mockRejectedValue(failure)

const result = getKimiCodeModels("token", { signal: controller.signal })
if (outcome === "success") await expect(result).resolves.toEqual({})
else await expect(result).rejects.toBe(failure)
expect(getEventListeners(controller.signal, "abort")).toHaveLength(0)
expect(vi.getTimerCount()).toBe(0)
controller.abort()
expect(transport.mock.calls[0][1]?.signal?.aborted).toBe(false)
})

it("aborts model discovery after its deadline", async () => {
vi.useFakeTimers()
vi.spyOn(globalThis, "fetch").mockImplementation((_input, init) => {
return new Promise((_resolve, reject) => {
init?.signal?.addEventListener("abort", () => reject(init.signal?.reason), { once: true })
})
})
const result = expect(getKimiCodeModels("token")).rejects.toThrow("timed out")
const controller = new AbortController()
const result = expect(getKimiCodeModels("token", { signal: controller.signal })).rejects.toThrow("timed out")

await vi.advanceTimersByTimeAsync(10_000)
await result
expect(vi.mocked(fetch).mock.calls[0][1]?.signal?.aborted).toBe(true)
expect(getEventListeners(controller.signal, "abort")).toHaveLength(0)
expect(vi.getTimerCount()).toBe(0)
})

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
import { getEventListeners } from "events"
import axios from "axios"
import { providerIdentifiers } from "@roo-code/types"

import { ModelRequestRegistry } from "../../../../core/webview/ModelRequestRegistry"
import { getModels, refreshModels } from "../modelCache"

vi.mock("../../../../services/zoo-code-auth", () => ({
getZooCodeBaseUrl: () => "https://example.test",
resolveZooGatewaySessionToken: (token?: string) => token,
}))
vi.mock("@roo-code/telemetry", () => ({
TelemetryService: { instance: { captureEvent: vi.fn(), isTelemetryEnabled: () => false } },
}))

// Keep the registry, cache dispatch, and auth-scoped fetchers real; only HTTP is replaced.
describe.each([getModels, refreshModels])("auth-scoped cancellation via %s", (loadModels) => {
describe.each([providerIdentifiers.zooGateway, providerIdentifiers.kimiCode])("%s", (provider) => {
let registry: ModelRequestRegistry
let signals: AbortSignal[]
let events: string[]

beforeEach(() => {
vi.useFakeTimers()
registry = new ModelRequestRegistry()
signals = []
events = []
vi.spyOn(console, "error").mockImplementation(() => {})
const pendingRequest = (signal?: AbortSignal | null) => {
expect(signal).toBeDefined()
if (!signal) throw new Error("HTTP request has no cancellation signal")
signals.push(signal)
const index = signals.length
events.push(`start:${index}`)
return new Promise<never>((_resolve, reject) => {
signal.addEventListener(
"abort",
() => {
events.push(`abort:${index}`)
reject(signal.reason)
},
{ once: true },
)
})
}
vi.spyOn(axios, "get").mockImplementation((_url, config) => {
expect(config?.timeout).toBe(15_000)
return pendingRequest(config?.signal instanceof AbortSignal ? config.signal : undefined)
})
vi.spyOn(globalThis, "fetch").mockImplementation((_url, init) => pendingRequest(init?.signal))
})

afterEach(() => {
registry.dispose()
vi.restoreAllMocks()
vi.useRealTimers()
})

it.each(["cancel", "replace"])("aborts HTTP before restarting after %s", async (action) => {
const ownerSignals: AbortSignal[] = []
const discover = async (signal: AbortSignal) => {
ownerSignals.push(signal)
await loadModels({ provider, apiKey: "synthetic-token", signal })
}
const first = registry.run("router", discover)
expect(events).toEqual(["start:1"])
if (action === "cancel") registry.cancel("router")
const second = registry.run("router", discover)
expect(events).toEqual(["start:1", "abort:1", "start:2"])
expect(signals[0].aborted).toBe(true)
expect(signals[1].aborted).toBe(false)
await first
registry.cancel("router")
await second
expect(events).toEqual(["start:1", "abort:1", "start:2", "abort:2"])
for (const signal of ownerSignals) expect(getEventListeners(signal, "abort")).toHaveLength(0)
expect(vi.getTimerCount()).toBe(0)
})

it("does not start HTTP for a pre-aborted caller", async () => {
const controller = new AbortController()
controller.abort()
const result = loadModels({ provider, apiKey: "synthetic-token", signal: controller.signal })
if (loadModels === getModels) await expect(result).rejects.toMatchObject({ name: "AbortError" })
else await expect(result).resolves.toEqual({})
expect(axios.get).not.toHaveBeenCalled()
expect(fetch).not.toHaveBeenCalled()
expect(vi.getTimerCount()).toBe(0)
})
})
})
Loading
Loading