import { afterAll, beforeAll, beforeEach, describe, expect, test, vi } from 'vitest'; import { promises as fs } from 'fs'; import os from 'os'; import path from 'path'; import type { NextRequest } from 'next/server'; const mocks = vi.hoisted(() => ({ generateTTS: vi.fn(), recordUsage: vi.fn(), serverConfigured: vi.fn(), })); vi.mock('@/lib/audio/tts-providers', () => ({ generateTTS: (...args: unknown[]) => mocks.generateTTS(...args), TTSRateLimitError: class TTSRateLimitError extends Error {}, })); vi.mock('@/lib/server/usage-storage', () => ({ recordGenerationUsage: (...args: unknown[]) => mocks.recordUsage(...args), })); vi.mock('@/lib/server/provider-config', () => ({ isServerConfiguredProvider: (..._args: unknown[]) => mocks.serverConfigured(..._args), isServerTTSProviderDisabled: () => false, resolveTTSApiKey: () => 'test-key', resolveTTSBaseUrl: () => undefined, resolveTTSModel: () => undefined, })); let cacheDir: string; beforeAll(async () => { cacheDir = await fs.mkdtemp(path.join(os.tmpdir(), 'qa-tts-cache-')); }); afterAll(async () => { await fs.rm(cacheDir, { recursive: true, force: true }); vi.unstubAllEnvs(); }); async function postTts(body: unknown, token = 'x') { const { POST } = await import('@/app/api/tts/route'); const request = new Request('http://localhost/api/tts', { method: 'POST', headers: { 'Content-Type': 'application/json', 'x-forwarded-for': token }, body: JSON.stringify(body), }); return POST(request as unknown as NextRequest); } describe('POST /api/tts', () => { beforeEach(() => { vi.resetModules(); vi.stubEnv('QA_TTS_PROVIDER', 'openai-tts'); vi.stubEnv('TTS_CACHE_DIR', cacheDir); vi.stubEnv('TTS_RATE_LIMIT_PER_MIN', '1000'); mocks.generateTTS.mockReset(); mocks.recordUsage.mockReset(); mocks.serverConfigured.mockReset().mockReturnValue(true); }); test('synthesizes on first request and serves the cache on the second', async () => { mocks.generateTTS.mockResolvedValue({ audio: Buffer.from('fake-mp3-bytes'), format: 'mp3' }); const first = await postTts({ text: '光合作用在哪里发生?' }); expect(first.status).toBe(200); expect(first.headers.get('content-type')).toBe('audio/mpeg'); expect(await first.text()).toBe('fake-mp3-bytes'); expect(mocks.generateTTS).toHaveBeenCalledTimes(1); expect(mocks.recordUsage).toHaveBeenCalledWith( expect.objectContaining({ kind: 'tts', quantity: 10 }), ); // Second identical request: cache hit, no synthesis. mocks.generateTTS.mockClear(); const second = await postTts({ text: '光合作用在哪里发生?' }); expect(second.status).toBe(200); expect(await second.text()).toBe('fake-mp3-bytes'); expect(mocks.generateTTS).not.toHaveBeenCalled(); }); test('different text misses the cache', async () => { mocks.generateTTS.mockResolvedValue({ audio: Buffer.from('a'), format: 'mp3' }); await postTts({ text: '问题甲' }); mocks.generateTTS.mockClear(); await postTts({ text: '问题乙' }); expect(mocks.generateTTS).toHaveBeenCalledTimes(1); }); test('rejects missing or oversized text', async () => { const missing = await postTts({}); expect(missing.status).toBe(400); const oversized = await postTts({ text: 'x'.repeat(2001) }); expect(oversized.status).toBe(400); }); test('503 when no TTS provider is configured', async () => { mocks.serverConfigured.mockReturnValue(false); vi.stubEnv('QA_TTS_PROVIDER', ''); vi.resetModules(); const response = await postTts({ text: '你好' }); expect(response.status).toBe(503); }); });