Skip to content

Commit 032ac70

Browse files
committed
fix(kimi-code): harden OAuth request handling
1 parent 2a2e4f3 commit 032ac70

8 files changed

Lines changed: 363 additions & 101 deletions

File tree

src/api/providers/__tests__/kimi-code.spec.ts

Lines changed: 57 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -1,37 +1,27 @@
11
import { buildApiHandler } from "../../index"
22
import { KimiCodeHandler } from "../kimi-code"
3-
import type { Mock } from "vitest"
4-
5-
vi.mock("../../../integrations/kimi-code/oauth", () => {
6-
const mockGetAccessToken = vi.fn().mockResolvedValue("oauth-token")
7-
const mockForceRefreshAccessToken = vi.fn().mockResolvedValue("refreshed-token")
8-
return {
9-
kimiCodeOAuthManager: {
10-
getAccessToken: mockGetAccessToken,
11-
forceRefreshAccessToken: mockForceRefreshAccessToken,
12-
},
13-
mockGetAccessToken,
14-
mockForceRefreshAccessToken,
15-
}
16-
})
173

18-
vi.mock("../fetchers/modelCache", () => {
19-
const mockGetModels = vi.fn().mockRejectedValue(new Error("offline"))
20-
return {
21-
getModels: mockGetModels,
22-
mockGetModels,
23-
}
24-
})
4+
const { mockGetAccessToken, mockForceRefreshAccessToken, mockGetModels } = vi.hoisted(() => ({
5+
mockGetAccessToken: vi.fn(),
6+
mockForceRefreshAccessToken: vi.fn(),
7+
mockGetModels: vi.fn(),
8+
}))
9+
10+
vi.mock("../../../integrations/kimi-code/oauth", () => ({
11+
kimiCodeOAuthManager: {
12+
getAccessToken: mockGetAccessToken,
13+
forceRefreshAccessToken: mockForceRefreshAccessToken,
14+
},
15+
}))
2516

26-
const { mockGetAccessToken, mockForceRefreshAccessToken } = await import("../../../integrations/kimi-code/oauth")
27-
const { mockGetModels } = await import("../fetchers/modelCache")
17+
vi.mock("../fetchers/modelCache", () => ({ getModels: mockGetModels }))
2818

2919
describe("KimiCodeHandler", () => {
3020
beforeEach(() => {
3121
vi.clearAllMocks()
32-
;(mockGetAccessToken as any).mockResolvedValue("oauth-token")
33-
;(mockForceRefreshAccessToken as any).mockResolvedValue("refreshed-token")
34-
;(mockGetModels as any).mockRejectedValue(new Error("offline"))
22+
mockGetAccessToken.mockResolvedValue("oauth-token")
23+
mockForceRefreshAccessToken.mockResolvedValue("refreshed-token")
24+
mockGetModels.mockRejectedValue(new Error("offline"))
3525
})
3626

3727
it("is dispatched separately from Moonshot and preserves an unknown selected model", () => {
@@ -65,7 +55,7 @@ describe("KimiCodeHandler", () => {
6555
} catch {
6656
// expected - mock is incomplete
6757
}
68-
expect(mockGetAccessToken as any).not.toHaveBeenCalled()
58+
expect(mockGetAccessToken).not.toHaveBeenCalled()
6959
})
7060

7161
it("uses OAuth token when auth method is oauth or not specified", async () => {
@@ -78,11 +68,11 @@ describe("KimiCodeHandler", () => {
7868
} catch {
7969
// expected - mock will fail
8070
}
81-
expect(mockGetAccessToken as any).toHaveBeenCalled()
71+
expect(mockGetAccessToken).toHaveBeenCalled()
8272
})
8373

8474
it("throws error when OAuth is required but no token available", async () => {
85-
;(mockGetAccessToken as any).mockResolvedValueOnce(null)
75+
mockGetAccessToken.mockResolvedValueOnce(null)
8676
const handler = new KimiCodeHandler({ kimiCodeAuthMethod: "oauth" })
8777
const gen = handler.createMessage("system", [{ role: "user", content: "test" }])
8878
await expect(async () => {
@@ -105,9 +95,7 @@ describe("KimiCodeHandler", () => {
10595
it("retries with forced refresh on 401 when using OAuth", async () => {
10696
const handler = new KimiCodeHandler({ kimiCodeAuthMethod: "oauth" })
10797
const fetchSpy = vi.spyOn(globalThis, "fetch")
108-
fetchSpy.mockResolvedValueOnce(
109-
new Response(null, { status: 401 }),
110-
)
98+
fetchSpy.mockResolvedValueOnce(new Response(null, { status: 401 }))
11199
fetchSpy.mockResolvedValueOnce(
112100
new Response(JSON.stringify({ choices: [{ message: { content: "ok" }, finish_reason: "stop" }] }), {
113101
status: 200,
@@ -121,25 +109,36 @@ describe("KimiCodeHandler", () => {
121109
} catch {
122110
// expected - mock is incomplete
123111
}
124-
expect(mockForceRefreshAccessToken as any).toHaveBeenCalled()
112+
expect(mockForceRefreshAccessToken).toHaveBeenCalledOnce()
113+
})
114+
115+
it("force-refreshes and retries exactly once after a non-streaming OAuth 401", async () => {
116+
const handler = new KimiCodeHandler({ kimiCodeAuthMethod: "oauth" })
117+
const unauthorized = Object.assign(new Error("Unauthorized"), { status: 401 })
118+
const createCompletion = vi
119+
.spyOn((handler as any).client.chat.completions, "create")
120+
.mockRejectedValueOnce(unauthorized)
121+
.mockResolvedValueOnce({ choices: [{ message: { content: "retried" } }] })
122+
123+
await expect(handler.completePrompt("test")).resolves.toBe("retried")
124+
expect(mockForceRefreshAccessToken).toHaveBeenCalledOnce()
125+
expect(createCompletion).toHaveBeenCalledTimes(2)
125126
})
126127

127128
it("does not retry on 401 when using API key auth", async () => {
128129
const handler = new KimiCodeHandler({ kimiCodeAuthMethod: "api-key", kimiCodeApiKey: "key" })
129-
const fetchSpy = vi.spyOn(globalThis, "fetch").mockResolvedValueOnce(
130-
new Response(null, { status: 401 }),
131-
)
130+
const fetchSpy = vi.spyOn(globalThis, "fetch").mockResolvedValueOnce(new Response(null, { status: 401 }))
132131
const gen = handler.createMessage("system", [{ role: "user", content: "test" }])
133132
await expect(async () => {
134133
for await (const chunk of gen) {
135134
// consume
136135
}
137136
}).rejects.toThrow()
138-
expect(mockForceRefreshAccessToken as any).not.toHaveBeenCalled()
137+
expect(mockForceRefreshAccessToken).not.toHaveBeenCalled()
139138
})
140139

141140
it("fetches models during prepareRequest", async () => {
142-
;(mockGetModels as any).mockResolvedValueOnce({ "test-model": { maxTokens: 1000 } })
141+
mockGetModels.mockResolvedValueOnce({ "test-model": { maxTokens: 1000 } })
143142
const handler = new KimiCodeHandler({ kimiCodeAuthMethod: "api-key", kimiCodeApiKey: "key" })
144143
const gen = handler.createMessage("system", [{ role: "user", content: "test" }])
145144
try {
@@ -149,11 +148,11 @@ describe("KimiCodeHandler", () => {
149148
} catch {
150149
// expected
151150
}
152-
expect(mockGetModels as any).toHaveBeenCalled()
151+
expect(mockGetModels).toHaveBeenCalled()
153152
})
154153

155154
it("continues when model discovery fails", async () => {
156-
;(mockGetModels as any).mockRejectedValueOnce(new Error("discovery failed"))
155+
mockGetModels.mockRejectedValueOnce(new Error("discovery failed"))
157156
const handler = new KimiCodeHandler({ kimiCodeAuthMethod: "api-key", kimiCodeApiKey: "key" })
158157
const gen = handler.createMessage("system", [{ role: "user", content: "test" }])
159158
try {
@@ -163,11 +162,27 @@ describe("KimiCodeHandler", () => {
163162
} catch {
164163
// expected - different error
165164
}
166-
expect(mockGetModels as any).toHaveBeenCalled()
165+
expect(mockGetModels).toHaveBeenCalled()
166+
})
167+
168+
it.each([
169+
["failure", () => Promise.reject(new Error("offline"))],
170+
["empty response", () => Promise.resolve({})],
171+
])("does not repeatedly block requests after model discovery %s", async (_case, discovery) => {
172+
mockGetModels.mockImplementationOnce(discovery)
173+
vi.spyOn(globalThis, "fetch").mockImplementation(
174+
async () => new Response(JSON.stringify({ choices: [{ message: { content: "ok" } }] }), { status: 200 }),
175+
)
176+
const handler = new KimiCodeHandler({ kimiCodeAuthMethod: "api-key", kimiCodeApiKey: "key" })
177+
178+
await handler.completePrompt("first")
179+
await handler.completePrompt("second")
180+
181+
expect(mockGetModels).toHaveBeenCalledOnce()
167182
})
168183

169184
it("uses discovered model info when available", async () => {
170-
;(mockGetModels as any).mockResolvedValueOnce({ "kimi-for-coding": { maxTokens: 8000, contextWindow: 128000 } })
185+
mockGetModels.mockResolvedValueOnce({ "kimi-for-coding": { maxTokens: 8000, contextWindow: 128000 } })
171186
const handler = new KimiCodeHandler({ kimiCodeAuthMethod: "api-key", kimiCodeApiKey: "key" })
172187
const gen = handler.createMessage("system", [{ role: "user", content: "test" }])
173188
try {

src/api/providers/__tests__/openai.spec.ts

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -905,6 +905,12 @@ describe("OpenAiHandler", () => {
905905
await expect(handler.completePrompt("Test prompt")).rejects.toThrow("OpenAI completion error: API Error")
906906
})
907907

908+
it("should preserve HTTP status when wrapping completion errors", async () => {
909+
mockCreate.mockRejectedValueOnce(Object.assign(new Error("Unauthorized"), { status: 401 }))
910+
911+
await expect(handler.completePrompt("Test prompt")).rejects.toMatchObject({ status: 401 })
912+
})
913+
908914
it("should handle empty response", async () => {
909915
mockCreate.mockImplementationOnce(() => ({
910916
choices: [{ message: { content: "" } }],

src/api/providers/fetchers/__tests__/kimi-code.spec.ts

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ import { getKimiCodeModels, mapKimiCodeModel } from "../kimi-code"
22

33
describe("Kimi Code model discovery", () => {
44
beforeEach(() => vi.restoreAllMocks())
5+
afterEach(() => vi.useRealTimers())
56

67
it("maps official model fields", () => {
78
expect(
@@ -77,4 +78,19 @@ describe("Kimi Code model discovery", () => {
7778
expect(models["model-a"].contextWindow).toBe(100000)
7879
expect(models["model-b"].supportsReasoningBinary).toBe(true)
7980
})
81+
82+
it("aborts model discovery after its deadline", async () => {
83+
vi.useFakeTimers()
84+
vi.spyOn(globalThis, "fetch").mockImplementation((_input, init) => {
85+
return new Promise((_resolve, reject) => {
86+
init?.signal?.addEventListener("abort", () => reject(init.signal?.reason), { once: true })
87+
})
88+
})
89+
const result = expect(getKimiCodeModels("token")).rejects.toThrow("timed out")
90+
91+
await vi.advanceTimersByTimeAsync(10_000)
92+
await result
93+
expect(vi.mocked(fetch).mock.calls[0][1]?.signal?.aborted).toBe(true)
94+
expect(vi.getTimerCount()).toBe(0)
95+
})
8096
})

src/api/providers/fetchers/kimi-code.ts

Lines changed: 21 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,8 @@ const kimiCodeModelSchema = z.object({
1212

1313
const kimiCodeModelsResponseSchema = z.object({ data: z.array(kimiCodeModelSchema) })
1414

15+
const KIMI_CODE_MODELS_TIMEOUT_MS = 10_000
16+
1517
export function mapKimiCodeModel(model: z.infer<typeof kimiCodeModelSchema>): ModelInfo {
1618
return {
1719
...kimiCodeDefaultModelInfo,
@@ -24,14 +26,24 @@ export function mapKimiCodeModel(model: z.infer<typeof kimiCodeModelSchema>): Mo
2426

2527
export async function getKimiCodeModels(apiKey?: string): Promise<ModelRecord> {
2628
if (!apiKey) throw new Error("Kimi Code authentication is required to fetch models")
27-
const response = await fetch(`${KIMI_CODE_BASE_URL}/models`, {
28-
headers: { Accept: "application/json", Authorization: `Bearer ${apiKey}` },
29-
})
30-
if (!response.ok) {
31-
const error = new Error(`Kimi Code models request failed: ${response.status} ${response.statusText}`)
32-
;(error as Error & { status?: number }).status = response.status
33-
throw error
29+
const controller = new AbortController()
30+
const timeout = setTimeout(
31+
() => controller.abort(new Error("Kimi Code models request timed out")),
32+
KIMI_CODE_MODELS_TIMEOUT_MS,
33+
)
34+
try {
35+
const response = await fetch(`${KIMI_CODE_BASE_URL}/models`, {
36+
headers: { Accept: "application/json", Authorization: `Bearer ${apiKey}` },
37+
signal: controller.signal,
38+
})
39+
if (!response.ok) {
40+
const error = new Error(`Kimi Code models request failed: ${response.status} ${response.statusText}`)
41+
;(error as Error & { status?: number }).status = response.status
42+
throw error
43+
}
44+
const parsed = kimiCodeModelsResponseSchema.parse(await response.json())
45+
return Object.fromEntries(parsed.data.map((model) => [model.id, mapKimiCodeModel(model)]))
46+
} finally {
47+
clearTimeout(timeout)
3448
}
35-
const parsed = kimiCodeModelsResponseSchema.parse(await response.json())
36-
return Object.fromEntries(parsed.data.map((model) => [model.id, mapKimiCodeModel(model)]))
3749
}

src/api/providers/kimi-code.ts

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ function getHttpStatus(error: unknown): number | undefined {
3131
export class KimiCodeHandler extends OpenAiHandler {
3232
private readonly kimiOptions: ApiHandlerOptions
3333
private models: ModelRecord = {}
34+
private modelDiscoveryAttempted = false
3435

3536
constructor(options: ApiHandlerOptions) {
3637
super({
@@ -59,11 +60,15 @@ export class KimiCodeHandler extends OpenAiHandler {
5960
private async prepareRequest(forceRefresh = false): Promise<void> {
6061
const accessToken = await this.resolveAccessToken(forceRefresh)
6162
this.client.apiKey = accessToken
62-
if (Object.keys(this.models).length === 0) {
63+
if (!this.modelDiscoveryAttempted) {
64+
this.modelDiscoveryAttempted = true
6365
try {
6466
this.models = await getModels({ provider: "kimi-code", apiKey: accessToken })
65-
} catch {
67+
} catch (error) {
6668
// Model discovery is best-effort; preserve the configured ID and fallback metadata.
69+
console.debug("[KimiCode] Model discovery failed; using fallback model metadata", {
70+
message: error instanceof Error ? error.message : String(error),
71+
})
6772
}
6873
}
6974
}

src/api/providers/openai.ts

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -323,7 +323,13 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
323323
return response.choices?.[0]?.message.content || ""
324324
} catch (error) {
325325
if (error instanceof Error) {
326-
throw new Error(`${this.providerName} completion error: ${error.message}`)
326+
const wrapped = new Error(`${this.providerName} completion error: ${error.message}`, { cause: error })
327+
const source = error as Error & { status?: number; errorDetails?: unknown; code?: unknown }
328+
const target = wrapped as Error & { status?: number; errorDetails?: unknown; code?: unknown }
329+
if (source.status !== undefined) target.status = source.status
330+
if (source.errorDetails !== undefined) target.errorDetails = source.errorDetails
331+
if (source.code !== undefined) target.code = source.code
332+
throw wrapped
327333
}
328334

329335
throw error

0 commit comments

Comments
 (0)