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
107 changes: 50 additions & 57 deletions src/api/providers/__tests__/gemini-handler.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,7 @@ vi.mock("@roo-code/telemetry", () => ({
}))

import { GeminiHandler } from "../gemini"
import type { ApiHandlerOptions } from "../../../shared/api"
import { providerIdentifiers } from "@roo-code/types/provider-identifiers"
import { makeApiHandlerOptions } from "../../../test-utils/api"

describe("GeminiHandler backend support", () => {
beforeEach(() => {
Expand All @@ -24,11 +23,7 @@ describe("GeminiHandler backend support", () => {
// URL context and grounding are mutually exclusive with function declarations
// in Gemini API, so createMessage only uses function declarations.
// URL context/grounding are only added in completePrompt.
const options = {
apiProvider: providerIdentifiers.gemini,
enableUrlContext: true,
enableGrounding: true,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand All @@ -41,11 +36,7 @@ describe("GeminiHandler backend support", () => {
})

it("completePrompt passes config overrides without tools when URL context and grounding disabled", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
enableUrlContext: false,
enableGrounding: false,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockResolvedValue({ text: "ok" })
// @ts-ignore access private client
Expand All @@ -56,12 +47,39 @@ describe("GeminiHandler backend support", () => {
expect(promptConfig.tools).toBeUndefined()
})

it("completePrompt should pass abort signal through to client via config.abortSignal", async () => {
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)

const controller = new AbortController()
const stub = vi.fn().mockResolvedValue({ text: "response" })
handler["client"].models.generateContent = stub

await handler.completePrompt("test prompt", { abortSignal: controller.signal })

expect(stub).toHaveBeenCalledWith(
expect.objectContaining({
config: expect.objectContaining({
abortSignal: controller.signal,
}),
}),
)
})

it("completePrompt should work without options (backward compatible)", async () => {
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)

const stub = vi.fn().mockResolvedValue({ text: "response" })
handler["client"].models.generateContent = stub

const result = await handler.completePrompt("test prompt")
expect(result).toBe("response")
})

describe("error scenarios", () => {
it("should handle grounding metadata extraction failure gracefully", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
enableGrounding: true,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)

const mockStream = async function* () {
Expand Down Expand Up @@ -93,10 +111,7 @@ describe("GeminiHandler backend support", () => {
})

it("should handle malformed grounding metadata", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
enableGrounding: true,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)

const mockStream = async function* () {
Expand Down Expand Up @@ -144,11 +159,7 @@ describe("GeminiHandler backend support", () => {
})

it("should handle API errors when tools are enabled", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
enableUrlContext: true,
enableGrounding: true,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)

const mockError = new Error("API rate limit exceeded")
Expand Down Expand Up @@ -192,9 +203,7 @@ describe("GeminiHandler backend support", () => {
]

it("should ignore allowedFunctionNames because Gemini rejects larger restriction lists", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand All @@ -213,9 +222,7 @@ describe("GeminiHandler backend support", () => {
})

it("should include all tools when allowedFunctionNames is provided", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand All @@ -236,9 +243,7 @@ describe("GeminiHandler backend support", () => {
})

it("should not pass large allowedFunctionNames lists to Gemini", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand Down Expand Up @@ -267,9 +272,7 @@ describe("GeminiHandler backend support", () => {
})

it("should not pass allowedFunctionNames even when history includes tool calls", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand Down Expand Up @@ -304,9 +307,7 @@ describe("GeminiHandler backend support", () => {
})

it("should fall back to tool_choice when allowedFunctionNames is provided", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand All @@ -327,9 +328,7 @@ describe("GeminiHandler backend support", () => {
})

it("should fall back to tool_choice when allowedFunctionNames is empty", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand All @@ -351,9 +350,7 @@ describe("GeminiHandler backend support", () => {
})

it("should not set toolConfig when allowedFunctionNames is undefined and no tool_choice", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand All @@ -374,9 +371,7 @@ describe("GeminiHandler backend support", () => {

describe("Gemini schema compatibility", () => {
it("should strip broad JSON Schema metadata from function declarations", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand Down Expand Up @@ -435,9 +430,7 @@ describe("GeminiHandler backend support", () => {
})

it("should collapse composition and type arrays in function declaration schemas", async () => {
const options = {
apiProvider: providerIdentifiers.gemini,
} as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand Down Expand Up @@ -496,7 +489,7 @@ describe("GeminiHandler backend support", () => {
})

it("should deep-merge allOf fragments instead of overwriting earlier properties", async () => {
const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand Down Expand Up @@ -541,7 +534,7 @@ describe("GeminiHandler backend support", () => {
})

it("should resolve $ref entries before dropping $defs", async () => {
const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand Down Expand Up @@ -590,7 +583,7 @@ describe("GeminiHandler backend support", () => {
})

it("should preserve top-level properties and required entries when allOf is also present", async () => {
const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand Down Expand Up @@ -632,7 +625,7 @@ describe("GeminiHandler backend support", () => {
})

it("should stop recursive $ref expansion before the sanitized schema becomes cyclic", async () => {
const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand Down Expand Up @@ -684,7 +677,7 @@ describe("GeminiHandler backend support", () => {
})

it("should preserve parameter names that collide with stripped schema keywords", async () => {
const options = { apiProvider: providerIdentifiers.gemini } as ApiHandlerOptions
const options = makeApiHandlerOptions()
const handler = new GeminiHandler(options)
const stub = vi.fn().mockReturnValue((async function* () {})())
// @ts-ignore access private client
Expand Down
Loading
Loading