Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions src/main/ipc/chatHandlers.ts
Original file line number Diff line number Diff line change
Expand Up @@ -106,17 +106,17 @@ export function registerChatHandlers(
}

// 3.2 RAG 增强:检索相关知识并注入上下文
// 只有在配置了 embedding connection 时才启用 RAG。
// 只要 embedding 后端可用就启用 RAG(远程 connection 或内置本地模型)。
// 以前这里用 `getEmbeddingClient()` 判断,它只认远程 connection,本地模型
// 直接返回 null —— 于是默认配置下 RAG 被静默关闭,回答无依据也无引用。
// 检索结果同时记录到消息上:以前检索失败只留一行日志,
// 于是「没有依据的回答」和「有依据的回答」在界面上完全无法区分。
let retrieval: RetrievalStatus = 'none'
let answerSources: AnswerSource[] = []
let answerCitations: Citation[] = []
let citationContexts: CitationContext[] = []
try {
const embeddingClient = await connectionManager.getEmbeddingClient()

if (embeddingClient) {
if (await knowledgeService.isEmbeddingAvailable()) {
const session = queries.getSessionById(sessionId)
if (session?.notebookId) {
const searchResults = await knowledgeService.search(session.notebookId, content, {
Expand Down Expand Up @@ -148,7 +148,7 @@ export function registerChatHandlers(
}
}
} else {
Logger.debug('ChatHandlers', 'RAG disabled: No embedding model configured')
Logger.debug('ChatHandlers', 'RAG disabled: no embedding backend available')
}
} catch (error) {
// RAG 失败不应该阻止对话
Expand Down
8 changes: 7 additions & 1 deletion src/main/models/ConnectionManager.ts
Original file line number Diff line number Diff line change
Expand Up @@ -77,7 +77,13 @@ export class ConnectionManager {
async getEmbeddingClient(): Promise<ModelClient | null> {
const connection = await connectionConfigManager.getConnection('embedding')
if (!connection) {
Logger.warn('ConnectionManager', 'No embedding model configured')
// 这不是错误。`EmbeddingService.resolveBackend()` 在拿不到远程 connection 时
// 会 fallback 到内置本地模型,也就是默认路径。用 debug,否则每个本地用户每次
// 索引/检索都会看到一条看起来像故障的 warning。
Logger.debug(
'ConnectionManager',
'No remote embedding connection configured; EmbeddingService will use the built-in local model'
)
return null
}
if (!protocolSupportsEmbedding(connection.protocol)) {
Expand Down
42 changes: 40 additions & 2 deletions src/main/services/EmbeddingService.ts
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ import type {
LocalEmbeddingModelInfo
} from '../../shared/types'
import { DEFAULT_EMBEDDING_SOURCES } from '../../shared/types'
import { ConnectionManager } from '../models/ConnectionManager'
import type { ConnectionManager } from '../models/ConnectionManager'
import Logger from '../../shared/utils/logger'
import { LocalEmbeddingBackend } from '../embedding/LocalEmbeddingBackend'
import type { TransformersModuleLoader } from '../embedding/LocalEmbeddingBackend'
Expand All @@ -36,6 +36,7 @@ import {
*/
export interface EmbeddingServiceConfig {
batchSize?: number // 远程批处理大小,默认 20
localBatchSize?: number // 本地批处理大小,默认 16
maxRetries?: number // 最大重试次数,默认 3
retryDelay?: number // 重试延迟(毫秒),默认 1000
rateLimit?: number // 请求间隔(毫秒),默认 100
Expand Down Expand Up @@ -78,6 +79,9 @@ export class EmbeddingService {
this.loadTransformers = options.loadTransformers
this.config = {
batchSize: options.batchSize ?? 20,
// 本地一次推理的输入量。multilingual-e5-small 按 512 token 截断,把整本书的
// chunk(几百条)一次喂给 pipeline 会同时炸内存和分词耗时;16 是保守起点。
localBatchSize: options.localBatchSize ?? 16,
maxRetries: options.maxRetries ?? 3,
retryDelay: options.retryDelay ?? 1000,
rateLimit: options.rateLimit ?? 100
Expand Down Expand Up @@ -153,7 +157,41 @@ export class EmbeddingService {
return await this.embedRemoteInBatches(backend, texts, purpose, onProgress)
}

return await this.withRetry(() => backend.embedBatch(texts, purpose, onProgress))
return await this.embedLocalInBatches(backend, texts, purpose, onProgress)
}

/**
* 本地推理分 batch。
*
* 不分 batch 时,整份文档的 chunk 会被一次性 tokenize 并送进一个 ONNX 前向:一本
* 几百 chunk 的书会长时间无响应、内存峰值失控,而且 `onProgress` 只在最后被调用一次,
* 界面上进度一直停在起点。分开之后峰值受控,进度按 batch 前进。
*/
private async embedLocalInBatches(
backend: EmbeddingBackend,
texts: string[],
purpose: EmbeddingPurpose,
onProgress?: (completed: number, total: number) => void
): Promise<BackendEmbeddingResult[]> {
const batches = this.chunk(texts, this.config.localBatchSize)
const results: BackendEmbeddingResult[] = []

Logger.info(
'EmbeddingService',
`Local embedding: ${texts.length} texts in ${batches.length} batches of up to ${this.config.localBatchSize}`
)

for (const batch of batches) {
const batchResults = await this.withRetry(() => backend.embedBatch(batch, purpose))
results.push(...batchResults)
onProgress?.(results.length, texts.length)

// 分词是同步 JS:一批的 tokenize 会把主进程事件循环按住。每个 batch 之间让出
// 一次,IPC(索引进度、其它窗口请求)才有机会被处理。
await new Promise<void>((resolve) => setImmediate(resolve))
}

return results
}

private async embedRemoteInBatches(
Expand Down
10 changes: 10 additions & 0 deletions src/main/services/KnowledgeService.ts
Original file line number Diff line number Diff line change
Expand Up @@ -740,6 +740,16 @@ export class KnowledgeService {
.run()
}

/**
* RAG 是否能跑:远程 embedding connection 已配置,或内置本地模型已安装。
*
* 不能用 `ConnectionManager.getEmbeddingClient()` 判断 —— 那个只看远程 connection,
* 本地模型时为 null,会把整条 RAG 路径关掉(本地是默认配置)。
*/
async isEmbeddingAvailable(): Promise<boolean> {
return await this.embeddingService.isAvailable()
}

/**
* 语义搜索。
*
Expand Down
46 changes: 46 additions & 0 deletions test/chatRagGate.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
import { test } from 'node:test'
import assert from 'node:assert/strict'
import { readFileSync } from 'node:fs'
import { join } from 'node:path'

/**
* chat 的 RAG 开关,以及「没有远程 connection」这条日志的语义。
*
* `ConnectionManager.getEmbeddingClient()` 只认远程 embedding connection;内置本地模型
* 是默认配置,它返回 null。以前 chat 用这个 null 当作「没有 embedding」的判据,于是默认
* 配置下整条 RAG 被静默关掉 —— 回答无依据、无引用,而索引其实完全正常。
*/

const read = (path: string): string => readFileSync(join(process.cwd(), path), 'utf8')

test('the built-in local fallback is not logged as a warning', () => {
const source = read('src/main/models/ConnectionManager.ts')

assert.ok(
!source.includes("'No embedding model configured'"),
'the local-fallback path must not keep the old warning text'
)
assert.ok(
source.includes('built-in local model'),
'the log should say the built-in local model is being used'
)
assert.match(source, /Logger\.debug\(/)
})

test('chat RAG is gated on embedding availability, not on a remote connection', () => {
const chat = read('src/main/ipc/chatHandlers.ts')

// Match the call, not the bare name: the explanatory comment mentions it.
assert.ok(
!/connectionManager\.getEmbeddingClient\(/.test(chat),
'getEmbeddingClient() is remote-only; using it as the gate disables RAG for local embeddings'
)
assert.ok(chat.includes('isEmbeddingAvailable'))

const knowledge = read('src/main/services/KnowledgeService.ts')
assert.ok(knowledge.includes('isEmbeddingAvailable'))
assert.ok(
knowledge.includes('embeddingService.isAvailable()'),
'isEmbeddingAvailable() must delegate to the backend availability check'
)
})
91 changes: 91 additions & 0 deletions test/embeddingBatching.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
import { test } from 'node:test'
import assert from 'node:assert/strict'
import {
EmbeddingService,
type EmbeddingServiceOptions
} from '../src/main/services/EmbeddingService.ts'
import type { ConnectionManager } from '../src/main/models/ConnectionManager.ts'
import type {
FeatureExtractionPipeline,
TransformersModuleLoader
} from '../src/main/embedding/LocalEmbeddingBackend.ts'

/**
* 本地 embedding 必须分批。
*
* 不分批时,整份文档的 chunk 会被一次性 tokenize 并送进一个 ONNX 前向:一本几百 chunk
* 的书会长时间无响应、内存峰值失控,而 `onProgress` 只在最后被调用一次,索引进度一直
* 停在起点。远程路径早就有 batch,本地这条没有。
*/

// SAFETY: EmbeddingService only reaches `getEmbeddingClient` and `getConnection` on its
// ConnectionManager, and this stub answers "no remote connection" for both so the service
// takes the built-in local path. Any new call site would fail loudly instead of silently
// reading a fake field.
const localOnlyConnectionManager = (): ConnectionManager =>
({
getEmbeddingClient: async () => null,
getConnection: async () => null
}) as unknown as ConnectionManager

/** Records the size of every array that reaches the ONNX pipeline. */
function localService(batchSizes: number[], options: Partial<EmbeddingServiceOptions> = {}) {
const loader: TransformersModuleLoader = async () => ({
env: { cacheDir: null, allowRemoteModels: false, allowLocalModels: true, useFSCache: true },
pipeline: async () => {
const extractor: FeatureExtractionPipeline = async (texts) => {
batchSizes.push(texts.length)
return { tolist: () => texts.map(() => [0.1, 0.2]) }
}
return extractor
}
})

return new EmbeddingService(localOnlyConnectionManager(), {
cacheDir: '/tmp/knownote-embedding-test',
loadTransformers: loader,
...options
})
}

test('local embeddings run in bounded batches, not one inference over everything', async () => {
const batchSizes: number[] = []
const service = localService(batchSizes, { localBatchSize: 16 })

const texts = Array.from({ length: 40 }, (_, index) => `chunk ${index}`)
const progress: Array<[number, number]> = []
const results = await service.embedBatch(texts, 'document', (completed, total) =>
progress.push([completed, total])
)

assert.deepEqual(batchSizes, [16, 16, 8], 'the whole array must not reach the pipeline at once')
assert.equal(results.length, 40)
// Progress advances per batch instead of jumping 0 -> total once at the end.
assert.deepEqual(progress, [
[16, 40],
[32, 40],
[40, 40]
])
})

test('the local batch size defaults to 16', async () => {
const batchSizes: number[] = []
const service = localService(batchSizes)

await service.embedBatch(
Array.from({ length: 33 }, (_, index) => `c${index}`),
'document'
)

assert.deepEqual(batchSizes, [16, 16, 1])
})

test('a single local embed still goes through the same pipeline', async () => {
const batchSizes: number[] = []
const service = localService(batchSizes)

const result = await service.embed('only', 'query')

assert.deepEqual(batchSizes, [1])
assert.equal(result.dimensions, 2)
})
12 changes: 8 additions & 4 deletions test/ts-resolve.mjs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,9 @@
* The main/renderer sources are written with extensionless relative imports and
* bundled by electron-vite, so `node --test` cannot resolve them directly. This
* hook adds `.ts` resolution for relative specifiers so tests can import
* main-process modules as-is, without rewriting the source conventions.
* main-process modules as-is, without rewriting the source conventions. Both
* shapes the sources use are covered: `./logger` -> `logger.ts` and
* `../../shared/types` -> `shared/types/index.ts`.
*/

import { registerHooks } from 'node:module'
Expand All @@ -19,9 +21,11 @@ registerHooks({
!/\.[a-zA-Z0-9]+$/.test(specifier) &&
context.parentURL?.startsWith('file:')
) {
const candidate = resolvePath(dirname(fileURLToPath(context.parentURL)), `${specifier}.ts`)
if (existsSync(candidate)) {
return nextResolve(pathToFileURL(candidate).href, context)
const base = resolvePath(dirname(fileURLToPath(context.parentURL)), specifier)
for (const candidate of [`${base}.ts`, resolvePath(base, 'index.ts')]) {
if (existsSync(candidate)) {
return nextResolve(pathToFileURL(candidate).href, context)
}
}
}
return nextResolve(specifier, context)
Expand Down
Loading