retrieval-models.spec.ts 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284
  1. import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
  2. import * as fs from 'node:fs'
  3. import * as os from 'node:os'
  4. import * as path from 'node:path'
  5. import { changeIndexState, forgetSceneBoundaries, searchFinalized, syncFinalizedIndex, withBookWrite, type EmbeddingProvider, type RerankingProvider, type SceneProvider } from '../src'
  6. import { scanFinalized } from '../src/retrieval/source'
  7. import { readSceneRecords, SCENE_CACHE_PATH, validateSceneEnds } from '../src/retrieval/scenes'
  8. import { removeSync } from '../src/repo/remove'
  9. let root: string
  10. beforeEach(() => { root = fs.mkdtempSync(path.join(os.tmpdir(), 'webnovel-retrieval-models-')) })
  11. afterEach(() => { removeSync(root) })
  12. function chapter(n: number, body = '甲说:救人。\n\n乙跳进河中。\n\n另一边,城门关闭。') {
  13. const file = path.join(root, '定稿', '卷01', String(n).padStart(4, '0') + '-章' + n + '.md')
  14. fs.mkdirSync(path.dirname(file), { recursive: true })
  15. fs.writeFileSync(file, '---\n版本: 1\n角色: 已定稿\n---\n' + body)
  16. return file
  17. }
  18. const embedding = (model = 'fixture'): EmbeddingProvider => {
  19. const run: EmbeddingProvider['embed'] = async inputs => inputs.map(() => [1, 0])
  20. return { metadata: { provider: 'test', model, revision: model, dimensions: 2 }, embed: run, embedBatch: run }
  21. }
  22. const scenes = (segment: SceneProvider['segment'] = async paragraphs => [paragraphs.length], concurrency = 1): SceneProvider => ({
  23. metadata: { provider: 'test', model: 'scene', revision: 'v1', concurrency }, segment,
  24. })
  25. const ranker = (rerank: RerankingProvider['rerank']): RerankingProvider => ({ metadata: { provider: 'test', model: 'rank', revision: 'v1', candidates: 40 }, rerank })
  26. function deferred<T>() {
  27. let resolve!: (value: T) => void
  28. const promise = new Promise<T>(done => { resolve = done })
  29. return { promise, resolve }
  30. }
  31. function deadlineClock() {
  32. const deadlines: Array<{ milliseconds: number; controller: AbortController }> = []
  33. const spy = vi.spyOn(AbortSignal, 'timeout').mockImplementation(milliseconds => {
  34. const controller = new AbortController()
  35. deadlines.push({ milliseconds, controller })
  36. return controller.signal
  37. })
  38. return { advanceTo: (milliseconds: number) => {
  39. for (const deadline of deadlines) if (deadline.milliseconds <= milliseconds) deadline.controller.abort()
  40. }, restore: () => spy.mockRestore() }
  41. }
  42. describe('场景边界与独立缓存', () => {
  43. it('并发识别通过串行短锁保存,不把同任务的进度写入误判为外部锁竞争', async () => {
  44. chapter(1); chapter(2, '另一章。\n\n另一段。'); chapter(3, '第三章独立正文。')
  45. const client = embedding(), sceneClient = scenes(async paragraphs => [paragraphs.length], 2)
  46. expect(await syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => sceneClient })).toMatchObject({
  47. ok: true, state: { phase: 'ready', scenes: { completed: 3, generated: 3, fallback: 0 } },
  48. })
  49. expect(Object.keys(readSceneRecords(root))).toHaveLength(3)
  50. })
  51. it.each([[], [1], [2, 2, 3], [3, 2], [0, 3], [1, 4], [1, NaN, 3], new Array(3)])('拒绝遗漏、重叠、越界与稀疏边界 %j', ends => {
  52. expect(() => validateSceneEnds(ends, 3)).toThrow()
  53. })
  54. it('按场景截取原文,长场景分块不跨场景,查询不请求场景模型', async () => {
  55. const body = Array.from({ length: 16 }, (_, i) => '段' + i + '山'.repeat(99)).join('\n\n')
  56. chapter(1, body)
  57. const segment = vi.fn<SceneProvider['segment']>(async () => [8, 16])
  58. const client = embedding()
  59. const scenesClient = scenes(segment)
  60. expect(await syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => scenesClient })).toMatchObject({ ok: true })
  61. const snapshot = await scanFinalized(root)
  62. const chunks = snapshot.documents[0]!.chunks
  63. expect(chunks.length).toBeGreaterThan(2)
  64. expect(new Set(chunks.map(chunk => chunk.scene?.id)).size).toBe(2)
  65. for (const chunk of chunks) {
  66. expect(chunk.text).toBe(body.slice(chunk.start, chunk.end))
  67. expect(chunk.startLine).toBeGreaterThanOrEqual(chunk.scene!.startLine)
  68. expect(chunk.endLine).toBeLessThanOrEqual(chunk.scene!.endLine)
  69. }
  70. const result = await searchFinalized(root, { query: '段8', provider: client })
  71. expect(result.ok && result.hits[0]?.scene).toBeDefined()
  72. expect(segment).toHaveBeenCalledTimes(1)
  73. })
  74. it('重启、改名、元数据、更换模型和向量重建复用场景,正文改动仅识别该章', async () => {
  75. const file = chapter(1)
  76. chapter(2, '别章。\n\n另一幕。')
  77. const segment = vi.fn<SceneProvider['segment']>(async p => [p.length])
  78. let client = embedding()
  79. let sceneClient = scenes(segment)
  80. const run = () => syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => sceneClient })
  81. expect((await run()).ok).toBe(true)
  82. expect(segment).toHaveBeenCalledTimes(2)
  83. fs.renameSync(file, path.join(path.dirname(file), '0001-改名.md'))
  84. const renamed = path.join(path.dirname(file), '0001-改名.md')
  85. fs.writeFileSync(renamed, fs.readFileSync(renamed, 'utf8').replace('版本: 1', '版本: 2'))
  86. client = embedding('new-embedding')
  87. sceneClient = { ...scenes(segment), metadata: { provider: 'another', model: 'new-scene-model', revision: 'v2' } }
  88. await changeIndexState(root, state => ({ ...state, rebuild: true }))
  89. expect((await run()).ok).toBe(true)
  90. expect(segment).toHaveBeenCalledTimes(2)
  91. fs.appendFileSync(renamed, '\n\n新增场景。')
  92. expect((await run()).ok).toBe(true)
  93. expect(segment).toHaveBeenCalledTimes(3)
  94. expect(fs.readFileSync(path.join(root, SCENE_CACHE_PATH), 'utf8')).not.toContain('新增场景')
  95. })
  96. it('取消保留成功章,未配合取消的模型不会发布迟到边界', async () => {
  97. chapter(1); chapter(2, '第二章。\n\n对话。')
  98. const entered = deferred<void>(), late = deferred<readonly number[]>()
  99. let calls = 0
  100. const abort = new AbortController()
  101. const client = embedding()
  102. const sceneClient = scenes(async p => { if (++calls === 2) { entered.resolve(); return late.promise }; return [p.length] })
  103. const running = syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => sceneClient, signal: abort.signal })
  104. await entered.promise
  105. abort.abort()
  106. expect(await running).toMatchObject({ ok: false, failure: { code: 'cancelled' } })
  107. expect(Object.keys(readSceneRecords(root))).toHaveLength(1)
  108. late.resolve([99])
  109. const resumed = vi.fn<SceneProvider['segment']>(async p => [p.length])
  110. const next = scenes(resumed)
  111. expect(await syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => next })).toMatchObject({ ok: true })
  112. expect(resumed).toHaveBeenCalledTimes(1)
  113. expect(Object.keys(readSceneRecords(root))).toHaveLength(2)
  114. })
  115. it('无效边界明确回退,普通更新不反复调用;显式重新识别仅清理指定章', async () => {
  116. chapter(1); chapter(2, '其他章。\n\n第二段。')
  117. const client = embedding()
  118. const invalid = vi.fn<SceneProvider['segment']>(async () => [999])
  119. const bad = scenes(invalid)
  120. expect(await syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => bad })).toMatchObject({ ok: true, state: { scenes: { fallback: 2 } } })
  121. expect((await syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => bad })).ok).toBe(true)
  122. expect(invalid).toHaveBeenCalledTimes(2)
  123. await forgetSceneBoundaries(root, 1)
  124. await changeIndexState(root, state => ({ ...state, paused: false }))
  125. const valid = vi.fn<SceneProvider['segment']>(async p => [p.length])
  126. const restored = scenes(valid)
  127. expect((await syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => restored })).ok).toBe(true)
  128. expect(valid).toHaveBeenCalledTimes(1)
  129. expect((await scanFinalized(root)).documents.map(doc => doc.sceneStatus)).toEqual(['ready', 'fallback'])
  130. })
  131. it('识别期间允许保存,源变化后不写入旧边界', async () => {
  132. const file = chapter(1)
  133. const entered = deferred<void>(), response = deferred<readonly number[]>()
  134. const client = embedding(), sceneClient = scenes(async () => { entered.resolve(); return response.promise })
  135. const run = syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => sceneClient })
  136. await entered.promise
  137. withBookWrite(root, () => fs.appendFileSync(file, '\n\n作者新写一段。'))
  138. response.resolve([3])
  139. expect(await run).toMatchObject({ ok: false, failure: { code: 'source-changed' } })
  140. expect(Object.keys(readSceneRecords(root))).toHaveLength(0)
  141. })
  142. it('损坏场景缓存被保留,不隐式清空后产生全书 LLM 请求', async () => {
  143. chapter(1)
  144. const client = embedding()
  145. await syncFinalizedIndex(root, { getProvider: () => client })
  146. const file = path.join(root, SCENE_CACHE_PATH)
  147. fs.writeFileSync(file, '{broken')
  148. const segment = vi.fn<SceneProvider['segment']>()
  149. const sceneClient = scenes(segment)
  150. expect(await syncFinalizedIndex(root, { getProvider: () => client, getSceneProvider: () => sceneClient })).toMatchObject({ ok: false, failure: { code: 'scene-cache-error' } })
  151. expect(segment).not.toHaveBeenCalled()
  152. expect(fs.readFileSync(file, 'utf8')).toBe('{broken')
  153. })
  154. })
  155. describe('查询阶段的可选重排序', () => {
  156. async function ready() {
  157. chapter(1, '测试:城门的争吵。'); chapter(2, '测试:河边救人的经过。')
  158. const client = embedding()
  159. expect((await syncFinalizedIndex(root, { getProvider: () => client })).ok).toBe(true)
  160. return client
  161. }
  162. it('遵循提供方的自定义等待时间,不在旧 15 秒处提前回退', async () => {
  163. const client = await ready()
  164. const clock = deadlineClock()
  165. const entered = deferred<void>(), reply = deferred<readonly number[]>()
  166. let requestSignal: AbortSignal | undefined
  167. const ranking: RerankingProvider = {
  168. metadata: { provider: 'test', model: 'rank', revision: 'v1', candidates: 40, timeoutMs: 45_000 },
  169. rerank: async (_query, _inputs, signal) => { requestSignal = signal; entered.resolve(); return reply.promise },
  170. }
  171. try {
  172. const running = searchFinalized(root, { query: '测试', limit: 1, provider: client, getReranker: () => ranking })
  173. await entered.promise
  174. clock.advanceTo(20_000)
  175. expect(requestSignal?.aborted).toBe(false)
  176. reply.resolve([0, 1])
  177. expect(await running).toMatchObject({ ok: true, reranking: { applied: true }, hits: [{ chapter: 2 }] })
  178. } finally { reply.resolve([0, 1]); clock.restore() }
  179. })
  180. it('达到自定义预算后仍能回退,即使提供方不配合取消', async () => {
  181. const client = await ready()
  182. const clock = deadlineClock()
  183. const entered = deferred<void>(), late = deferred<readonly number[]>()
  184. const ranking: RerankingProvider = { metadata: { provider: 'test', model: 'rank', revision: 'v1', timeoutMs: 45_000 },
  185. rerank: async () => { entered.resolve(); return late.promise } }
  186. try {
  187. const running = searchFinalized(root, { query: '测试', limit: 1, provider: client, getReranker: () => ranking })
  188. await entered.promise
  189. clock.advanceTo(45_000)
  190. expect(await running).toMatchObject({ ok: true, reranking: { applied: false, reason: '重排等待超过 45 秒,本次使用原排序' }, hits: [{ chapter: 1 }] })
  191. } finally { late.resolve([0, 1]); clock.restore() }
  192. })
  193. it.each([undefined, 0, NaN, Infinity, 120_001, 1.5])('旧提供方或非法超时声明 %s 使用 30 秒有界默认', async timeoutMs => {
  194. const client = await ready()
  195. const clock = deadlineClock()
  196. const entered = deferred<void>(), late = deferred<readonly number[]>()
  197. let requestSignal: AbortSignal | undefined
  198. const ranking: RerankingProvider = { metadata: { provider: 'test', model: 'rank', revision: 'v1', timeoutMs },
  199. rerank: async (_query, _inputs, signal) => { requestSignal = signal; entered.resolve(); return late.promise } }
  200. try {
  201. const running = searchFinalized(root, { query: '测试', provider: client, getReranker: () => ranking })
  202. await entered.promise
  203. clock.advanceTo(20_000)
  204. expect(requestSignal?.aborted).toBe(false)
  205. clock.advanceTo(30_000)
  206. expect(await running).toMatchObject({ ok: true, reranking: { applied: false, reason: '重排等待超过 30 秒,本次使用原排序' } })
  207. } finally { late.resolve([0, 1]); clock.restore() }
  208. })
  209. it('等待中超时声明发生变化时不发布旧结果', async () => {
  210. const client = await ready()
  211. const entered = deferred<void>(), reply = deferred<readonly number[]>()
  212. const metadata = { provider: 'test', model: 'rank', revision: 'v1', timeoutMs: 45_000 }
  213. const ranking: RerankingProvider = { metadata, rerank: async () => { entered.resolve(); return reply.promise } }
  214. const running = searchFinalized(root, { query: '测试', provider: client, getReranker: () => ranking })
  215. await entered.promise
  216. metadata.timeoutMs = 60_000
  217. reply.resolve([0, 1])
  218. expect(await running).toMatchObject({ ok: false, code: 'provider-changed' })
  219. })
  220. it('从最终 limit 之外重排候选,并保留原文定位', async () => {
  221. const client = await ready()
  222. const rerank = vi.fn<RerankingProvider['rerank']>(async (_q, inputs) => inputs.map(input => input.text.includes('救人') ? 1 : 0))
  223. const ranking = ranker(rerank)
  224. const result = await searchFinalized(root, { query: '测试', limit: 1, provider: client, getReranker: () => ranking })
  225. expect(result).toMatchObject({ ok: true, reranking: { applied: true, candidates: 2 }, hits: [{ chapter: 2 }] })
  226. expect(rerank.mock.calls[0]![1]).toHaveLength(2)
  227. if (result.ok) expect(fs.readFileSync(result.hits[0]!.absolutePath, 'utf8')).toContain(result.hits[0]!.snippet)
  228. })
  229. it.each([[], [1, NaN], new Array(2)])('无效评分降级到原排序 %j', scores => {
  230. return (async () => {
  231. const client = await ready()
  232. const ranking = ranker(async () => scores)
  233. expect(await searchFinalized(root, { query: '测试', limit: 1, provider: client, getReranker: () => ranking })).toMatchObject({ ok: true, reranking: { applied: false }, hits: [{ chapter: 1 }] })
  234. })()
  235. })
  236. it('明确关键词模式和重排关闭时不额外调用模型', async () => {
  237. const client = await ready()
  238. const embed = vi.spyOn(client, 'embed')
  239. const rerank = vi.fn<RerankingProvider['rerank']>(async () => [0, 1])
  240. const ranking = ranker(rerank)
  241. expect(await searchFinalized(root, { query: '测试', mode: 'keyword', provider: client, getReranker: () => ranking })).toMatchObject({ ok: true, mode: 'keyword' })
  242. expect(embed).not.toHaveBeenCalled()
  243. expect(rerank).not.toHaveBeenCalled()
  244. await searchFinalized(root, { query: '测试', provider: client, rerank: false, getReranker: () => ranking })
  245. expect(rerank).not.toHaveBeenCalled()
  246. })
  247. it('重排网络等待不持写锁,源变化后拒收结果', async () => {
  248. const client = await ready()
  249. const entered = deferred<void>(), reply = deferred<readonly number[]>()
  250. const ranking = ranker(async () => { entered.resolve(); return reply.promise })
  251. const result = searchFinalized(root, { query: '测试', provider: client, getReranker: () => ranking })
  252. await entered.promise
  253. withBookWrite(root, () => fs.appendFileSync(path.join(root, '定稿/卷01/0001-章1.md'), '作者改动'))
  254. reply.resolve([0, 1])
  255. expect(await result).toMatchObject({ ok: false, code: 'source-changed' })
  256. })
  257. it('配置变化和取消均不会发布迟到重排', async () => {
  258. const client = await ready()
  259. const entered = deferred<void>(), reply = deferred<readonly number[]>()
  260. let ranking = ranker(async () => { entered.resolve(); return reply.promise })
  261. const result = searchFinalized(root, { query: '测试', provider: client, getReranker: () => ranking })
  262. await entered.promise
  263. ranking = ranker(async () => [1, 0])
  264. reply.resolve([0, 1])
  265. expect(await result).toMatchObject({ ok: false, code: 'provider-changed' })
  266. const abort = new AbortController(), waiting = deferred<void>(), late = deferred<readonly number[]>()
  267. ranking = ranker(async () => { waiting.resolve(); return late.promise })
  268. const cancelled = searchFinalized(root, { query: '测试', provider: client, getReranker: () => ranking, signal: abort.signal })
  269. await waiting.promise
  270. abort.abort()
  271. expect(await cancelled).toMatchObject({ ok: false, code: 'cancelled' })
  272. late.resolve([0, 1])
  273. })
  274. })