AIAssistant.utils.test.ts 5.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159
  1. import type { UIMessage } from 'ai'
  2. import { describe, expect, test } from 'vitest'
  3. import {
  4. hasPendingToolApproval,
  5. isReadOnlySelect,
  6. resolvePendingToolApprovalsAsDenied,
  7. } from './AIAssistant.utils'
  8. const createMessageWithPart = (
  9. part: UIMessage['parts'][number],
  10. role: UIMessage['role'] = 'assistant'
  11. ) =>
  12. [
  13. {
  14. id: `${role}-msg-1`,
  15. role,
  16. parts: [part],
  17. },
  18. ] as UIMessage[]
  19. const pendingSqlApprovalPart = {
  20. type: 'tool-execute_sql',
  21. toolCallId: 'call-1',
  22. state: 'approval-requested',
  23. input: { sql: 'select 1', label: 'Test query' },
  24. approval: { id: 'approval-1' },
  25. } satisfies UIMessage['parts'][number]
  26. describe('AIAssistant.utils.ts:isReadOnlySelect', () => {
  27. test('Should return true for SQL that only contains SELECT operation', () => {
  28. const sql = 'select * from countries where id > 100 order by id asc;'
  29. const result = isReadOnlySelect(sql)
  30. expect(result).toBe(true)
  31. })
  32. test('Should return false for SQL that contains INSERT operation', () => {
  33. const sql = `insert into countries (id, name) values (1, 'hello');`
  34. const result = isReadOnlySelect(sql)
  35. expect(result).toBe(false)
  36. })
  37. test('Should return false for SQL that contains UPDATE operation', () => {
  38. const sql = `update countries set name = 'hello' where id = 2;`
  39. const result = isReadOnlySelect(sql)
  40. expect(result).toBe(false)
  41. })
  42. test('Should return false for SQL that contains DELETE operation', () => {
  43. const sql = `delete from countries where id = 2;`
  44. const result = isReadOnlySelect(sql)
  45. expect(result).toBe(false)
  46. })
  47. test('Should return false for SQL that contains ALTER operation', () => {
  48. const sql = `alter table countries drop column id if exists;`
  49. const result = isReadOnlySelect(sql)
  50. expect(result).toBe(false)
  51. })
  52. test('Should return false for SQL that contains DROP operation', () => {
  53. const sql = `drop table if exists countries;`
  54. const result = isReadOnlySelect(sql)
  55. expect(result).toBe(false)
  56. })
  57. test('Should return false for SQL that contains CREATE operation', () => {
  58. const sql = `create schema test_schema;`
  59. const result = isReadOnlySelect(sql)
  60. expect(result).toBe(false)
  61. })
  62. test('Should return false for SQL that contains REPLACE operation', () => {
  63. const sql = `create or replace view test_view as select * from countries where id > 500;`
  64. const result = isReadOnlySelect(sql)
  65. expect(result).toBe(false)
  66. })
  67. test('Should return false for SQL that calls a function not whitelisted', () => {
  68. const sql = `select create_new_user();`
  69. const result = isReadOnlySelect(sql)
  70. expect(result).toBe(false)
  71. })
  72. test('Should return true for SQL that calls a function that is whitelisted', () => {
  73. const sql = `select count(select * from countries);`
  74. const result = isReadOnlySelect(sql)
  75. expect(result).toBe(true)
  76. })
  77. test('Should return false for SQL that contains a write operation with a read operation', () => {
  78. const sql1 = `select count(select * from countries); create schema joshen;`
  79. const result1 = isReadOnlySelect(sql1)
  80. expect(result1).toBe(false)
  81. const sql2 = `create schema joshen; select count(select * from countries);`
  82. const result2 = isReadOnlySelect(sql2)
  83. expect(result2).toBe(false)
  84. })
  85. })
  86. describe('AIAssistant.utils.ts:hasPendingToolApproval', () => {
  87. test('Should return true when an assistant message has a pending tool approval', () => {
  88. const messages = createMessageWithPart(pendingSqlApprovalPart)
  89. expect(hasPendingToolApproval(messages)).toBe(true)
  90. })
  91. test('Should return false when approval has already been answered', () => {
  92. const messages = createMessageWithPart({
  93. type: 'tool-execute_sql',
  94. toolCallId: 'call-1',
  95. state: 'approval-responded',
  96. input: { sql: 'select 1', label: 'Test query' },
  97. approval: { id: 'approval-1', approved: true },
  98. })
  99. expect(hasPendingToolApproval(messages)).toBe(false)
  100. })
  101. test('Should ignore non-assistant messages', () => {
  102. const messages = createMessageWithPart(pendingSqlApprovalPart, 'user')
  103. expect(hasPendingToolApproval(messages)).toBe(false)
  104. })
  105. test('Should return true for dynamic tool approvals', () => {
  106. const messages = createMessageWithPart({
  107. type: 'dynamic-tool',
  108. toolName: 'execute_sql',
  109. toolCallId: 'call-1',
  110. state: 'approval-requested',
  111. input: { sql: 'select 1', label: 'Test query' },
  112. approval: { id: 'approval-1' },
  113. })
  114. expect(hasPendingToolApproval(messages)).toBe(true)
  115. })
  116. })
  117. describe('AIAssistant.utils.ts:resolvePendingToolApprovalsAsDenied', () => {
  118. test('Should convert pending approvals into denied outputs', () => {
  119. const messages = createMessageWithPart(pendingSqlApprovalPart)
  120. const resolvedMessages = resolvePendingToolApprovalsAsDenied(messages)
  121. const [part] = resolvedMessages[0].parts
  122. expect(part).toMatchObject({
  123. state: 'output-denied',
  124. approval: {
  125. id: 'approval-1',
  126. approved: false,
  127. reason: 'Skipped because the user sent a follow-up message.',
  128. },
  129. })
  130. })
  131. test('Should leave completed approvals unchanged', () => {
  132. const messages = createMessageWithPart({
  133. type: 'tool-execute_sql',
  134. toolCallId: 'call-1',
  135. state: 'output-available',
  136. input: { sql: 'select 1', label: 'Test query' },
  137. output: [],
  138. approval: { id: 'approval-1', approved: true },
  139. })
  140. expect(resolvePendingToolApprovalsAsDenied(messages)).toEqual(messages)
  141. })
  142. })