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
8 changes: 8 additions & 0 deletions electron.vite.config.ts
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,14 @@ export default defineConfig({
plugins: [copyMigrationsPlugin()],
build: {
rollupOptions: {
// Two entries: the Electron main process and the `worker_threads` entry that
// runs local ONNX inference off the main process (#176). Both land in
// out/main, so `WorkerEmbeddingBackend` can resolve the worker as a sibling
// of the main bundle.
input: {
index: resolve('src/main/index.ts'),
embeddingWorker: resolve('src/main/embedding/embeddingWorker.ts')
},
// Only things rollup CANNOT inline may stay external. Every entry in
// package.json `dependencies` is already externalized automatically
// (electron-vite's build.externalizeDeps defaults to true); the two
Expand Down
4 changes: 4 additions & 0 deletions src/main/config/defaults.ts
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,10 @@ export const defaultSettings: AppSettings = {
language: 'en-US',
autoLaunch: false,
hasCompletedOnboarding: false,
// Off by default: it spends provider tokens, so it is opt-in. Two continuations
// is the default bound when it is on (#179).
autoContinueOnTruncation: false,
maxAutoContinueAttempts: 2,
prompts: {
mindMap: {
'zh-CN': `你是知识结构分析专家,负责从笔记本内容中提炼核心知识结构。
Expand Down
6 changes: 6 additions & 0 deletions src/main/config/settingsManager.ts
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,12 @@ function mergeSettings(stored: Partial<AppSettings>): AppSettings {
language: stored.language ?? defaultSettings.language,
autoLaunch: stored.autoLaunch ?? defaultSettings.autoLaunch,
hasCompletedOnboarding: stored.hasCompletedOnboarding ?? defaultSettings.hasCompletedOnboarding,
// An older store has neither key; the defaults are what keeps them meaningful
// rather than `undefined` reaching the chat handler (#179).
autoContinueOnTruncation:
stored.autoContinueOnTruncation ?? defaultSettings.autoContinueOnTruncation,
maxAutoContinueAttempts:
stored.maxAutoContinueAttempts ?? defaultSettings.maxAutoContinueAttempts,
prompts: {
mindMap: {
'zh-CN': stored.prompts?.mindMap?.['zh-CN'] ?? defaultSettings.prompts!.mindMap!['zh-CN'],
Expand Down
203 changes: 203 additions & 0 deletions src/main/embedding/WorkerEmbeddingBackend.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,203 @@
/**
* WorkerEmbeddingBackend(#176)
*
* `EmbeddingBackend` 的本地实现,但推理不在主线程:它把文本交给 `embeddingWorker`
* 线程,ONNX 的 tokenize + 前向在那里跑。主进程事件循环只剩消息收发。
*
* 批次与进度由 worker 负责(见 embeddingWorkerProtocol.ts),所以这个后端声明
* `batchesInternally = true`,EmbeddingService 不会在外面再切一遍。
*/

import { Worker } from 'worker_threads'
import { join } from 'path'
import type { EmbeddingPurpose, EmbeddingSpace } from '../../shared/types'
import Logger from '../../shared/utils/logger'
import { getLocalSpace } from './space'
import type { BackendEmbeddingResult, EmbeddingBackend } from './types'
import type {
EmbeddingWorkerInit,
EmbeddingWorkerResponse,
EmbeddingWorkerSnapshot
} from './embeddingWorkerProtocol'

export interface WorkerEmbeddingBackendOptions {
/** userData/models */
cacheDir: string
/** 一次 ONNX 前向的文本数 */
batchSize: number
/** 单个 batch 失败后的重试上限(worker 内部执行) */
maxRetries?: number
/** 重试基础退避(毫秒) */
retryDelay?: number
/** 覆盖 worker 入口,测试注入用 */
workerPath?: string
/** pipeline 首次加载进度(0-1) */
onLoadProgress?: (progress: number) => void
}

interface PendingRequest {
resolve: (results: BackendEmbeddingResult[]) => void
reject: (error: Error) => void
onProgress?: (completed: number, total: number) => void
}

export class WorkerEmbeddingBackend implements EmbeddingBackend {
readonly kind = 'local' as const
/** 批处理在 worker 内完成,外层不要再切 */
readonly batchesInternally = true

private readonly cacheDir: string
private readonly batchSize: number
private readonly maxRetries: number
private readonly retryDelay: number
private readonly explicitWorkerPath?: string
private readonly onLoadProgress?: (progress: number) => void
private readonly space: EmbeddingSpace = getLocalSpace()

private worker: Worker | null = null
private nextRequestId = 1
private readonly pending = new Map<number, PendingRequest>()

constructor(options: WorkerEmbeddingBackendOptions) {
this.cacheDir = options.cacheDir
this.batchSize = Math.max(1, options.batchSize)
this.maxRetries = Math.max(1, options.maxRetries ?? 3)
this.retryDelay = options.retryDelay ?? 1000
this.explicitWorkerPath = options.workerPath
this.onLoadProgress = options.onLoadProgress
}

getSpace(): EmbeddingSpace {
return this.space
}

async isReady(): Promise<boolean> {
return this.worker !== null
}

async embed(text: string, purpose: EmbeddingPurpose): Promise<BackendEmbeddingResult> {
const [result] = await this.embedBatch([text], purpose)
return result
}

async embedBatch(
texts: string[],
purpose: EmbeddingPurpose,
onProgress?: (completed: number, total: number) => void
): Promise<BackendEmbeddingResult[]> {
if (texts.length === 0) {
return []
}

const worker = this.ensureWorker()
const id = this.nextRequestId++

return await new Promise<BackendEmbeddingResult[]>((resolve, reject) => {
this.pending.set(id, { resolve, reject, onProgress })
worker.postMessage({ type: 'embed', id, texts, purpose })
})
}

/**
* 结束 worker 线程并释放 ONNX session。之后再调用会重新拉起一个 worker。
*/
async dispose(): Promise<void> {
const worker = this.worker
this.worker = null
this.failPending(new Error('Embedding worker disposed'))

if (!worker) return

const disposed = new Promise<void>((resolve) => {
worker.once('exit', () => resolve())
worker.once('error', () => resolve())
})
try {
worker.postMessage({ type: 'dispose', id: this.nextRequestId++ })
} catch {
// worker 已经不在了,terminate 兜底
}
worker.removeAllListeners('message')
await Promise.race([disposed, worker.terminate().then(() => undefined)])
}

private resolveWorkerPath(): string {
if (this.explicitWorkerPath) {
return this.explicitWorkerPath
}
// 主进程 bundle 与 worker 产物同在 out/main(electron.vite.config.ts 的第二个入口)。
return join(__dirname, 'embeddingWorker.js')
}

private ensureWorker(): Worker {
if (this.worker) {
return this.worker
}

const init: EmbeddingWorkerInit = {
cacheDir: this.cacheDir,
batchSize: this.batchSize,
maxRetries: this.maxRetries,
retryDelay: this.retryDelay
}
const worker = new Worker(this.resolveWorkerPath(), { workerData: init })

worker.on('message', (message: EmbeddingWorkerResponse) => this.handleMessage(message))
worker.on('error', (error) => {
Logger.error('WorkerEmbeddingBackend', 'Embedding worker failed:', error)
this.worker = null
this.failPending(error instanceof Error ? error : new Error(String(error)))
})
worker.on('exit', (code) => {
if (this.worker === worker) {
this.worker = null
}
if (code !== 0) {
this.failPending(new Error(`Embedding worker exited with code ${code}`))
}
})

this.worker = worker
return worker
}

private handleMessage(message: EmbeddingWorkerResponse): void {
if (message.type === 'load-progress') {
this.onLoadProgress?.(message.progress)
return
}

const pending = this.pending.get(message.id)
if (!pending) return

if (message.type === 'progress') {
pending.onProgress?.(message.completed, message.total)
return
}

this.pending.delete(message.id)

if (message.type === 'result') {
pending.resolve(message.results.map(toBackendResult))
} else if (message.type === 'error') {
pending.reject(new Error(message.message))
} else if (message.type === 'disposed') {
// dispose 请求没有调用方等待;正常路径由 dispose() 自己结算
}
}

private failPending(error: Error): void {
for (const pending of this.pending.values()) {
pending.reject(error)
}
this.pending.clear()
}
}

function toBackendResult(snapshot: EmbeddingWorkerSnapshot): BackendEmbeddingResult {
return {
embedding: snapshot.embedding,
model: snapshot.model,
dimensions: snapshot.dimensions
}
}
121 changes: 121 additions & 0 deletions src/main/embedding/embeddingWorker.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,121 @@
/**
* embeddingWorker(#176)
*
* 本地 embedding 的 tokenize + ONNX 前向在 `worker_threads` 里跑,不在 Electron 主进程。
* 主进程只负责发文本、收向量、转发进度;事件循环因此不再被一个 batch 的原生调用按住。
*
* 批次与进度都在这里:worker 收到整批文本后自己按 `batchSize` 切开,每完成一批回一次
* `progress`。这样「批大小」只有一处定义,主进程不会把一本书一次性推给 ONNX。
*/

import { parentPort, workerData } from 'worker_threads'
import { LocalEmbeddingBackend } from './LocalEmbeddingBackend'
import Logger from '../../shared/utils/logger'
import type {
EmbeddingWorkerInit,
EmbeddingWorkerRequest,
EmbeddingWorkerSnapshot
} from './embeddingWorkerProtocol'

if (!parentPort) {
throw new Error('embeddingWorker must be started as a worker thread')
}

/** 非空断言集中在这里:上面的检查之后,闭包里用到的就是这个常量。 */
const port = parentPort

const init = workerData as EmbeddingWorkerInit

const backend = new LocalEmbeddingBackend({
cacheDir: init.cacheDir,
onLoadProgress: (progress) => port.postMessage({ type: 'load-progress', progress })
})

function chunk<T>(items: T[], size: number): T[][] {
const batches: T[][] = []
for (let index = 0; index < items.length; index += size) {
batches.push(items.slice(index, index + size))
}
return batches
}

async function handleEmbed(
id: number,
texts: string[],
purpose: 'query' | 'document'
): Promise<void> {
const batches = chunk(texts, Math.max(1, init.batchSize))
const results: EmbeddingWorkerSnapshot[] = []

for (const batch of batches) {
const batchResults = await embedBatchWithRetry(batch, purpose, id)
for (const result of batchResults) {
results.push({
embedding: result.embedding,
model: result.model,
dimensions: result.dimensions
})
}
port.postMessage({
type: 'progress',
id,
completed: results.length,
total: texts.length
})
}

port.postMessage({ type: 'result', id, results })
}

/**
* 逐 batch 重试(而不是整批重来):一次 book-sized 导入里某个 batch 抖一下,不应该让
* 已经算好的几百条向量白做。重试策略跟 EmbeddingService 的远程路径同源。
*/
async function embedBatchWithRetry(
batch: string[],
purpose: 'query' | 'document',
id: number
): Promise<Array<{ embedding: Float32Array; model: string; dimensions: number }>> {
let lastError: Error = new Error('Embedding failed')

for (let attempt = 1; attempt <= Math.max(1, init.maxRetries); attempt++) {
try {
return await backend.embedBatch(batch, purpose)
} catch (error) {
lastError = error instanceof Error ? error : new Error(String(error))
Logger.warn(
'EmbeddingWorker',
`attempt ${attempt} failed for request ${id}: ${lastError.message}`
)
if (attempt < Math.max(1, init.maxRetries)) {
await new Promise((resolve) => setTimeout(resolve, init.retryDelay * 2 ** (attempt - 1)))
}
}
}

throw lastError
}

async function handle(message: EmbeddingWorkerRequest): Promise<void> {
if (message.type === 'dispose') {
await backend.dispose()
port.postMessage({ type: 'disposed', id: message.id })
return
}

try {
await handleEmbed(message.id, message.texts, message.purpose)
} catch (error) {
port.postMessage({
type: 'error',
id: message.id,
message: error instanceof Error ? error.message : String(error)
})
}
}

// 串行处理:ONNX session 不并发,排队也避免多个 inference 同时抢内存。
let queue: Promise<void> = Promise.resolve()
port.on('message', (message: EmbeddingWorkerRequest) => {
queue = queue.then(() => handle(message)).catch(() => undefined)
})
Loading
Loading