generate-v4.test.ts 2.2 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586
  1. import { safeSql } from '@supabase/pg-meta'
  2. import { UIMessage } from 'ai'
  3. import { expect, test, vi } from 'vitest'
  4. import generateV4 from '../../pages/api/ai/sql/generate-v4'
  5. import { sanitizeMessagePart } from '@/lib/ai/tools/tool-sanitizer'
  6. vi.mock('@/lib/ai/tools/tool-sanitizer', () => ({
  7. sanitizeMessagePart: vi.fn((part) => part),
  8. }))
  9. test('generateV4 calls the tool sanitizer', async () => {
  10. const mockReq = {
  11. method: 'POST',
  12. headers: {
  13. authorization: 'Bearer test-token',
  14. },
  15. body: {
  16. messages: [
  17. {
  18. id: 'test-msg-id',
  19. role: 'assistant',
  20. parts: [
  21. {
  22. type: 'tool-execute_sql',
  23. state: 'output-available',
  24. toolCallId: 'test-tool-call-id',
  25. input: { sql: safeSql`SELECT * FROM users` },
  26. output: [{ id: 1, name: 'test-output' }],
  27. },
  28. ],
  29. },
  30. ] satisfies UIMessage[],
  31. projectRef: 'test-project',
  32. connectionString: 'test-connection',
  33. orgSlug: 'test-org',
  34. },
  35. on: vi.fn(),
  36. }
  37. const mockRes = {
  38. status: vi.fn(() => mockRes),
  39. json: vi.fn(() => mockRes),
  40. setHeader: vi.fn(() => mockRes),
  41. }
  42. vi.mock('@/lib/ai/ai-details', () => ({
  43. getOrgAIDetails: vi.fn().mockResolvedValue({
  44. aiOptInLevel: 'schema_and_log_and_data',
  45. hasAccessToAdvanceModel: true,
  46. }),
  47. getProjectAIDetails: vi.fn().mockResolvedValue({
  48. region: 'us-east-1',
  49. isSensitive: false,
  50. }),
  51. }))
  52. vi.mock('@/lib/ai/model', () => ({
  53. getModel: vi.fn().mockResolvedValue({
  54. modelParams: { model: {} },
  55. systemProviderOptions: {},
  56. }),
  57. }))
  58. vi.mock('@/data/sql/execute-sql-query', () => ({
  59. executeSql: vi.fn().mockResolvedValue({ result: [] }),
  60. }))
  61. vi.mock('@/lib/ai/tools', () => ({
  62. getTools: vi.fn().mockResolvedValue({}),
  63. }))
  64. vi.mock('ai', async () => {
  65. const actual = await vi.importActual('ai')
  66. return {
  67. ...actual,
  68. streamText: vi.fn().mockReturnValue({
  69. pipeUIMessageStreamToResponse: vi.fn(),
  70. }),
  71. }
  72. })
  73. await generateV4(mockReq as any, mockRes as any)
  74. expect(sanitizeMessagePart).toHaveBeenCalled()
  75. })