From 3634f49a48e6984d01d5b0bf001f37f445e6cf40 Mon Sep 17 00:00:00 2001 From: zhenhao <33592619+huangzhenhao90@users.noreply.github.com> Date: Mon, 7 Sep 2026 15:07:22 +0800 Subject: [PATCH] fix(ai): send stable session headers for OpenCode Go --- .../services/agent-group-report-service.ts | 6 +- src/main/services/ai-provider-service.ts | 24 ++- .../src/hooks/useGroupReportGeneration.ts | 10 +- src/shared/ai-provider.ts | 2 + .../unit/ai-provider-opencode-session.test.ts | 157 ++++++++++++++++++ 5 files changed, 192 insertions(+), 7 deletions(-) create mode 100644 tests/unit/ai-provider-opencode-session.test.ts diff --git a/src/main/services/agent-group-report-service.ts b/src/main/services/agent-group-report-service.ts index 19100da..b0ff672 100644 --- a/src/main/services/agent-group-report-service.ts +++ b/src/main/services/agent-group-report-service.ts @@ -1,3 +1,4 @@ +import { randomUUID } from 'crypto' import type { Contact, Message } from '../../shared/types' import type { GroupDailyReport, GroupReportMetadata } from '../../shared/group-report' import { exportGroupReport } from '../group-report-service' @@ -120,12 +121,13 @@ export async function generateAgentGroupReport( const input = await buildGroupReportInput(messages, contact as Contact, true, 'full') const runtime = aiProvider.getRuntimeConfig() + const sessionId = randomUUID() const ai = await aiProvider.chat( [ { role: 'system', content: GROUP_REPORT_SYSTEM_PROMPT }, { role: 'user', content: input.prompt } ], - { timeoutMs: Math.max(30, Math.min(1800, request.timeoutSeconds || 300)) * 1000 } + { sessionId, timeoutMs: Math.max(30, Math.min(1800, request.timeoutSeconds || 300)) * 1000 } ) if (!ai.success || !ai.data) { return { @@ -153,7 +155,7 @@ export async function generateAgentGroupReport( { role: 'system', content: GROUP_REPORT_JSON_REPAIR_SYSTEM_PROMPT }, { role: 'user', content: ai.data } ], - { timeoutMs: Math.max(30, Math.min(1800, request.timeoutSeconds || 300)) * 1000 } + { sessionId, timeoutMs: Math.max(30, Math.min(1800, request.timeoutSeconds || 300)) * 1000 } ) if (!repaired.success || !repaired.data) { const cause = parseError instanceof Error ? parseError.message : String(parseError) diff --git a/src/main/services/ai-provider-service.ts b/src/main/services/ai-provider-service.ts index 48bcf8b..6a98e24 100644 --- a/src/main/services/ai-provider-service.ts +++ b/src/main/services/ai-provider-service.ts @@ -1,4 +1,5 @@ import { app } from 'electron' +import { randomUUID } from 'crypto' import fs from 'fs-extra' import path from 'path' import type { @@ -399,7 +400,8 @@ export class AIProviderService { messages, testing, signal, - onDelta + onDelta, + options?.sessionId ) } @@ -427,14 +429,15 @@ export class AIProviderService { onDelta?: AIChatDeltaHandler ): Promise { const provider = deepSeekProvider(options.baseURL, options.model) - return requestOpenAICompatible( + return requestProvider( provider, options.apiKey || '', options.model || provider.defaultModel, messages, false, signal, - onDelta + onDelta, + options?.sessionId ) } @@ -640,8 +643,21 @@ function requestProvider( messages: AIMessage[], testing = false, signal?: AbortSignal, - onDelta?: AIChatDeltaHandler + onDelta?: AIChatDeltaHandler, + sessionId?: string ): Promise { + // Match the endpoint, not the editable provider name or model name. + const url = new URL(provider.baseUrl) + if (url.hostname === 'opencode.ai' && /^\/zen\/go(?:\/|$)/.test(url.pathname)) { + const extraHeaders = Object.fromEntries( + Object.entries(provider.advanced.extraHeaders).filter( + ([name]) => name.toLowerCase() !== 'x-opencode-session' + ) + ) + extraHeaders['x-opencode-session'] = sessionId?.trim() || randomUUID() + if (!hasHeader(extraHeaders, 'user-agent')) extraHeaders['user-agent'] = 'TraceMemo' + provider = { ...provider, advanced: { ...provider.advanced, extraHeaders } } + } if (provider.type === 'anthropic-messages') { return requestAnthropic(provider, apiKey, model, messages, testing, signal) } diff --git a/src/renderer/src/hooks/useGroupReportGeneration.ts b/src/renderer/src/hooks/useGroupReportGeneration.ts index cbb5711..014555c 100644 --- a/src/renderer/src/hooks/useGroupReportGeneration.ts +++ b/src/renderer/src/hooks/useGroupReportGeneration.ts @@ -91,6 +91,7 @@ interface UseGroupReportGenerationArgs { } interface PreparedReportContext { + sessionId: string input: Awaited> startedAt: number logs: ReportGenerationLog[] @@ -503,6 +504,7 @@ export function useGroupReportGeneration({ const result = await trackStep(`AI 生成(${selectedModel.model})`, () => withTimeout( window.api.aiChat(aiMessages, { + sessionId: context.sessionId, providerId: selectedModel.providerId, modelId: selectedModel.model, timeoutMs: reportTimeoutSeconds * 1000 @@ -556,6 +558,7 @@ export function useGroupReportGeneration({ const repairResult = await trackStep(`AI 修复 JSON(${selectedModel.model})`, () => withTimeout( window.api.aiChat(repairMessages, { + sessionId: context.sessionId, providerId: selectedModel.providerId, modelId: selectedModel.model, timeoutMs: reportTimeoutSeconds * 1000 @@ -785,7 +788,12 @@ export function useGroupReportGeneration({ imageInsightsInjectedIntoPrompt: input.imageInsightSummary.succeeded > 0 && input.prompt.includes('AI 图片识别摘要:') }) - const context: PreparedReportContext = { input, startedAt: startGenerateTime, logs } + const context: PreparedReportContext = { + sessionId: crypto.randomUUID(), + input, + startedAt: startGenerateTime, + logs + } preparedContextRef.current = context if (input.imageInsightSummary.failed > 0) { setPreparationProgress({ diff --git a/src/shared/ai-provider.ts b/src/shared/ai-provider.ts index 00b59f6..eb55722 100644 --- a/src/shared/ai-provider.ts +++ b/src/shared/ai-provider.ts @@ -117,6 +117,8 @@ export interface LegacyAIConfig { } export interface AIChatRequestOptions { + /** Stable ID shared by requests and retries belonging to one conversation or task. */ + sessionId?: string providerId?: string modelId?: string timeoutMs?: number diff --git a/tests/unit/ai-provider-opencode-session.test.ts b/tests/unit/ai-provider-opencode-session.test.ts new file mode 100644 index 0000000..af14a91 --- /dev/null +++ b/tests/unit/ai-provider-opencode-session.test.ts @@ -0,0 +1,157 @@ +import { mkdtempSync, rmSync } from 'fs' +import { tmpdir } from 'os' +import { join } from 'path' +import { afterAll, afterEach, describe, expect, it, vi } from 'vitest' +import type { AIProviderConfig } from '../../src/shared/ai-provider' + +const root = mkdtempSync(join(tmpdir(), 'tracememo-opencode-session-')) +vi.mock('electron', () => ({ + app: { getPath: () => root }, + safeStorage: { + isEncryptionAvailable: () => true, + encryptString: (value: string) => Buffer.from(value), + decryptString: (value: Buffer) => value.toString('utf8') + } +})) +import { AIProviderService } from '../../src/main/services/ai-provider-service' + +const messages = [{ role: 'user', content: 'hello' }] +const uuid = /^[0-9a-f]{8}-(?:[0-9a-f]{4}-){3}[0-9a-f]{12}$/i +function setup(overrides: Partial = {}): { + service: AIProviderService + fetchMock: ReturnType +} { + const service = new AIProviderService() + const provider: AIProviderConfig = { + id: 'custom-name', + name: 'My model', + type: 'openai-compatible', + baseUrl: 'https://opencode.ai/zen/go/v1', + auth: { type: 'bearer' }, + apiKey: 'test-key', + models: [ + { + id: 'test-model', + name: 'Test', + capabilities: { chat: true, vision: true, ocr: true, longContext: true } + } + ], + defaultModel: 'test-model', + advanced: { timeoutMs: 1000, extraHeaders: { 'x-custom': 'kept' } }, + ...overrides + } + expect(service.save(provider).success).toBe(true) + service.setDefault(provider.id) + const fetchMock = vi.fn().mockImplementation( + async () => + new Response( + JSON.stringify({ + choices: [{ message: { content: 'ok' }, finish_reason: 'stop' }], + output_text: 'ok', + status: 'completed', + content: [{ type: 'text', text: 'ok' }], + stop_reason: 'end_turn' + }), + { headers: { 'content-type': 'application/json' } } + ) + ) + vi.stubGlobal('fetch', fetchMock) + return { service, fetchMock } +} +function headers(fetchMock: ReturnType, index = 0): Headers { + return new Headers(fetchMock.mock.calls[index][1].headers) +} +afterEach(() => vi.unstubAllGlobals()) +afterAll(() => rmSync(root, { recursive: true, force: true })) + +describe('OpenCode Go session headers', () => { + it('keeps a task ID across concurrent calls and retries without sharing it with another task', async () => { + const { service, fetchMock } = setup() + await Promise.all( + ['task-a', 'task-b', 'task-a'].map((sessionId) => service.chat(messages, { sessionId })) + ) + expect([0, 1, 2].map((i) => headers(fetchMock, i).get('x-opencode-session'))).toEqual([ + 'task-a', + 'task-b', + 'task-a' + ]) + expect(headers(fetchMock).get('authorization')).toBe('Bearer test-key') + expect(headers(fetchMock).get('x-custom')).toBe('kept') + expect(headers(fetchMock).get('user-agent')).toBe('TraceMemo') + expect( + service.list().providers.find((p) => p.id === 'custom-name')?.advanced.extraHeaders + ).toEqual({ 'x-custom': 'kept' }) + }) + + it.each(['chat-completions', 'responses', 'anthropic'] as const)( + 'covers %s inference and connection tests', + async (protocol) => { + const { service, fetchMock } = setup({ + type: protocol === 'anthropic' ? 'anthropic-messages' : 'openai-compatible', + advanced: { + timeoutMs: 1000, + extraHeaders: {}, + apiProtocol: protocol === 'responses' ? 'responses' : 'chat-completions' + } + }) + expect(await service.chat(messages)).toMatchObject({ success: true, data: 'ok' }) + expect(await service.test('custom-name')).toMatchObject({ success: true }) + expect(headers(fetchMock).get('x-opencode-session')).toMatch(uuid) + expect(headers(fetchMock, 1).get('x-opencode-session')).toMatch(uuid) + expect(headers(fetchMock, 1).get('x-opencode-session')).not.toBe( + headers(fetchMock).get('x-opencode-session') + ) + } + ) + + it('covers image analysis and legacy caller options', async () => { + const { service, fetchMock } = setup() + expect( + await service.analyzeImage( + [ + { + role: 'user', + content: [{ type: 'image', dataUrl: 'data:image/png;base64,iVBORw0KGgo=' }] + } + ], + { sessionId: 'image-task' } + ) + ).toMatchObject({ + success: true + }) + expect( + await service.chat(messages, { + apiKey: 'legacy-key', + baseURL: 'https://opencode.ai/zen/go/v1', + model: 'test-model', + sessionId: 'legacy-task' + }) + ).toMatchObject({ success: true }) + expect(headers(fetchMock).get('x-opencode-session')).toBe('image-task') + expect(headers(fetchMock, 1).get('x-opencode-session')).toBe('legacy-task') + }) + + it('replaces stale static session headers case insensitively and preserves a custom user agent', async () => { + const { service, fetchMock } = setup({ + advanced: { + timeoutMs: 1000, + extraHeaders: { 'X-OpenCode-Session': 'static', 'User-Agent': 'CustomTraceMemo' } + } + }) + await service.chat(messages, { sessionId: 'new-task' }) + expect(headers(fetchMock).get('x-opencode-session')).toBe('new-task') + expect(headers(fetchMock).get('user-agent')).toBe('CustomTraceMemo') + }) + + it.each([ + 'https://api.openai.com/v1', + 'https://opencode.ai/zen/v1', + 'https://opencode.ai.example.com/zen/go/v1', + 'https://opencode.ai/zen/gopher' + ])('does not inject Go headers for %s', async (baseUrl) => { + const { service, fetchMock } = setup({ baseUrl }) + await service.chat(messages, { sessionId: 'private-task' }) + expect(headers(fetchMock).has('x-opencode-session')).toBe(false) + expect(headers(fetchMock).has('user-agent')).toBe(false) + }) +})