diff --git a/src/chat-view.ts b/src/chat-view.ts index 8a294fd..2e4403d 100755 --- a/src/chat-view.ts +++ b/src/chat-view.ts @@ -72,6 +72,15 @@ export class ChatView extends ItemView { undefined, newSettings.cacheConfig ); + void this.ollamaClient.initializeCache().catch(() => { + new Notice( + 'Semantic cache unavailable: could not connect to ChromaDB. Check the ChromaDB URL in settings.' + ); + }); + } + + public async clearCache(): Promise { + await this.ollamaClient.clearCache(); } getViewType(): string { @@ -83,7 +92,13 @@ export class ChatView extends ItemView { } async onOpen(): Promise { - await this.ollamaClient.initializeCache(); + try { + await this.ollamaClient.initializeCache(); + } catch { + new Notice( + 'Semantic cache unavailable: could not connect to ChromaDB. Check the ChromaDB URL in settings.' + ); + } this.render(); this.removeEventListeners(); // Clean up any existing listeners before reattaching this.setupEventListeners(); @@ -93,7 +108,7 @@ export class ChatView extends ItemView { this.updateSettings(newSettings); } - async onClose(): Promise { + onClose(): Promise { this.ollamaClient.cancelStream(); this.removeEventListeners(); this.cleanupStreamingResources(); @@ -268,7 +283,7 @@ export class ChatView extends ItemView { if (streamingMessage && !this.lastMessageEl) { this.lastMessageEl = (this.chatContainer ?? this.contentEl).createEl('div', { cls: `ollama-message assistant`, - }) as HTMLElement; + }); this.lastMessageEl.setAttribute('data-msg-id', streamingMessage.id); } if (this.lastMessageEl) { @@ -314,14 +329,11 @@ export class ChatView extends ItemView { return [ systemMessage, - ...this.messages.map( - (m) => - ({ - role: m.role, - content: m.content, - tool_calls: m.tool_calls, - }) as OllamaMessage - ), + ...this.messages.map((m) => ({ + role: m.role, + content: m.content, + tool_calls: m.tool_calls, + })), userMessageWithContext, ]; } diff --git a/src/constants.ts b/src/constants.ts index 85626bf..b425f6f 100644 --- a/src/constants.ts +++ b/src/constants.ts @@ -11,5 +11,6 @@ export const DEFAULT_SETTINGS = { similarityThreshold: 0.85, collectionName: 'ollama_semantic_cache', embeddingModel: 'nomic-embed-text', + chromaUrl: 'http://localhost:8000', }, }; diff --git a/src/main.ts b/src/main.ts index d6fbee9..e38e247 100755 --- a/src/main.ts +++ b/src/main.ts @@ -49,7 +49,11 @@ export default class OllamaPlugin extends Plugin { const data = (await this.loadData()) as Partial | null; if (data) { Logger.debug('Loading saved settings', 'settings'); - Object.assign(this.settings, data); + this.settings = { + ...DEFAULT_SETTINGS, + ...data, + cacheConfig: { ...DEFAULT_SETTINGS.cacheConfig, ...data.cacheConfig }, + }; } } catch (error) { ErrorHandler.handleError(error, 'settings load'); @@ -89,6 +93,17 @@ export default class OllamaPlugin extends Plugin { } }); } + + public async clearSemanticCache(): Promise { + const leaves = this.app.workspace.getLeavesOfType('ollama-chat-view'); + for (const leaf of leaves) { + const view = leaf.view; + if (view instanceof ChatView) { + await view.clearCache(); + return; + } + } + } } class OllamaSettingTab extends PluginSettingTab { @@ -152,6 +167,61 @@ class OllamaSettingTab extends PluginSettingTab { this.plugin.notifyChatViews(); }) ); + + new Setting(container) + .setName('ChromaDB URL') + .setDesc('URL of your ChromaDB instance (used for semantic cache)') + .addText((text) => + text.setValue(this.plugin.settings.cacheConfig.chromaUrl).onChange(async (value) => { + this.plugin.settings.cacheConfig.chromaUrl = value; + await this.plugin.saveSettings(); + this.plugin.notifyChatViews(); + }) + ); + + new Setting(container) + .setName('Cache Embedding Model') + .setDesc('Ollama model used to generate embeddings for the semantic cache') + .addText((text) => + text.setValue(this.plugin.settings.cacheConfig.embeddingModel).onChange(async (value) => { + this.plugin.settings.cacheConfig.embeddingModel = value; + await this.plugin.saveSettings(); + this.plugin.notifyChatViews(); + }) + ); + + new Setting(container) + .setName('Cache Similarity Threshold') + .setDesc( + 'Minimum cosine similarity (0–1) for a cache hit. Higher values require closer matches.' + ) + .addText((text) => + text + .setValue(String(this.plugin.settings.cacheConfig.similarityThreshold)) + .onChange(async (value) => { + const parsed = parseFloat(value); + if (!isNaN(parsed) && parsed >= 0 && parsed <= 1) { + this.plugin.settings.cacheConfig.similarityThreshold = parsed; + await this.plugin.saveSettings(); + } else { + new Notice('Similarity threshold must be a number between 0 and 1.'); + } + }) + ); + + new Setting(container) + .setName('Clear Semantic Cache') + .setDesc('Delete all cached responses from ChromaDB') + .addButton((button) => + button.setButtonText('Clear Cache').onClick(async () => { + try { + await this.plugin.clearSemanticCache(); + new Notice('Semantic cache cleared.'); + } catch { + new Notice('Failed to clear semantic cache. Is ChromaDB running?'); + } + }) + ); } hide(): void { diff --git a/src/ollama-client.ts b/src/ollama-client.ts index 625efc8..d49eb03 100644 --- a/src/ollama-client.ts +++ b/src/ollama-client.ts @@ -1,9 +1,9 @@ // src/ollama-client.ts import type { OllamaMessage, OllamaTool } from './types'; -import { ApiError } from './types'; +import { ApiError, CacheConfig } from './types'; import { Logger } from './utils'; -import { SemanticCacheService, CacheConfig } from './semantic-cache'; +import { SemanticCacheService } from './semantic-cache'; interface OllamaChatResponse { message?: Partial; @@ -33,6 +33,12 @@ export class OllamaClient { } } + async clearCache(): Promise { + if (this.cacheService) { + await this.cacheService.clearCache(); + } + } + cancelStream(): void { if (this.currentStreamController) { this.currentStreamController.abort(); @@ -50,7 +56,7 @@ export class OllamaClient { return; } - const lastUserMsg = [...messages].reverse().find((m) => m.role === 'user'); + const lastUserMsg = messages.findLast((m) => m.role === 'user'); if (lastUserMsg && this.cacheService) { const cached = await this.cacheService.getCache(lastUserMsg.content); if (cached) { @@ -59,16 +65,16 @@ export class OllamaClient { } } - const chunks: OllamaMessage[] = []; - for await (const chunk of this.streamChatWithRetry(messages, tools, 0)) { - chunks.push(chunk); - yield chunk; - } - - // Populate cache in background after successful stream if (this.cacheService && lastUserMsg) { + const chunks: OllamaMessage[] = []; + for await (const chunk of this.streamChatWithRetry(messages, tools, 0)) { + chunks.push(chunk); + yield chunk; + } const fullContent = chunks.map((c) => c.content).join(''); void this.cacheService.setCache(lastUserMsg.content, fullContent); + } else { + yield* this.streamChatWithRetry(messages, tools, 0); } } } @@ -90,7 +96,7 @@ export class OllamaClient { return this.chatWithRetry(messages, tools, 0); } - const lastUserMsg = [...messages].reverse().find((m) => m.role === 'user'); + const lastUserMsg = messages.findLast((m) => m.role === 'user'); if (lastUserMsg && this.cacheService) { const cached = await this.cacheService.getCache(lastUserMsg.content); if (cached) { @@ -328,7 +334,9 @@ export class OllamaClient { private throwIfOllamaError(parsed: Record): void { if (parsed.error) { - throw new Error(`Ollama error: ${String(parsed.error)}`); + const errorMsg = + typeof parsed.error === 'string' ? parsed.error : JSON.stringify(parsed.error); + throw new Error(`Ollama error: ${errorMsg}`); } } diff --git a/src/semantic-cache.ts b/src/semantic-cache.ts index b95f286..540e7be 100644 --- a/src/semantic-cache.ts +++ b/src/semantic-cache.ts @@ -1,22 +1,24 @@ // src/semantic-cache.ts -import { ChromaClient } from 'chromadb'; +import { ChromaClient, Collection, IncludeEnum } from 'chromadb'; import { Logger } from './utils'; import { CacheConfig } from './types'; +export { CacheConfig } from './types'; + export class SemanticCacheService { private client: ChromaClient; - private collection: ReturnType | null = null; + private collection: Collection | null = null; private config: CacheConfig; private ollamaURL: string; constructor(ollamaURL: string, config: CacheConfig) { this.ollamaURL = ollamaURL.replace(/\/+$/, ''); this.config = config; - this.client = new ChromaClient({ path: 'http://localhost:8000' }); + this.client = new ChromaClient({ path: config.chromaUrl }); } - async initialize() { + async initialize(): Promise { if (!this.config.enabled) return; try { @@ -27,9 +29,32 @@ export class SemanticCacheService { Logger.info(`Semantic cache initialized: ${this.config.collectionName}`, 'semantic-cache'); } catch (error) { Logger.error(`Failed to initialize semantic cache: ${String(error)}`, 'semantic-cache'); + throw error; } } + async clearCache(): Promise { + await this.client.deleteCollection({ name: this.config.collectionName }); + this.collection = null; + Logger.info('Semantic cache cleared', 'semantic-cache'); + + try { + await this.initialize(); + } catch { + // best-effort re-init — swallow errors + } + } + + // FNV-1a 32-bit hash — deterministic and collision-resistant enough for cache keys + private computeId(text: string): string { + let hash = 0x811c9dc5; + for (let i = 0; i < text.length; i++) { + hash ^= text.charCodeAt(i); + hash = Math.imul(hash, 0x01000193) >>> 0; + } + return hash.toString(16).padStart(8, '0'); + } + private async getEmbedding(text: string): Promise { try { const response = await fetch(`${this.ollamaURL}/api/embeddings`, { @@ -45,7 +70,7 @@ export class SemanticCacheService { throw new Error(`Embedding failed with status ${response.status}`); } - const data = await response.json(); + const data = (await response.json()) as { embedding: number[] }; return data.embedding; } catch (error) { Logger.warn(`Failed to generate embedding: ${String(error)}`, 'semantic-cache'); @@ -65,7 +90,7 @@ export class SemanticCacheService { const results = await this.collection.query({ queryEmbeddings: [embedding], nResults: 1, - include: ['metadatas', 'distances'], + include: [IncludeEnum.Metadatas, IncludeEnum.Distances], }); // Cosine distance = 1 - cosine_similarity @@ -76,7 +101,8 @@ export class SemanticCacheService { results.distances[0][0] < 1 - this.config.similarityThreshold ) { Logger.debug('Semantic cache hit', 'semantic-cache'); - return results.metadatas?.[0]?.[0]?.fullResponse ?? null; + const fullResponse = results.metadatas?.[0]?.[0]?.fullResponse; + return typeof fullResponse === 'string' ? fullResponse : null; } } catch (error) { Logger.warn(`Cache lookup failed: ${String(error)}`, 'semantic-cache'); @@ -94,8 +120,9 @@ export class SemanticCacheService { const embedding = await this.getEmbedding(prompt); if (!embedding.length) return; - await this.collection.add({ - ids: [crypto.randomUUID()], + const id = this.computeId(prompt); + await this.collection.upsert({ + ids: [id], embeddings: [embedding], metadatas: [{ fullResponse: response }], }); diff --git a/src/types.ts b/src/types.ts index f82099f..3b9f35d 100644 --- a/src/types.ts +++ b/src/types.ts @@ -95,6 +95,7 @@ export interface CacheConfig { similarityThreshold: number; collectionName: string; embeddingModel: string; + chromaUrl: string; } export interface PluginSettings { diff --git a/tests/chat-view.test.ts b/tests/chat-view.test.ts index b69df42..60af677 100755 --- a/tests/chat-view.test.ts +++ b/tests/chat-view.test.ts @@ -32,6 +32,13 @@ const mockSettings: PluginSettings = { vaultSearchLimit: 3, maxMessageHistory: 50, lastIndexTime: 0, + cacheConfig: { + enabled: false, + similarityThreshold: 0.9, + collectionName: 'test-cache', + embeddingModel: 'nomic-embed-text', + chromaUrl: 'http://localhost:8000', + }, }; describe('ChatView', () => { diff --git a/tests/ollama-client-cache.test.ts b/tests/ollama-client-cache.test.ts index efc4107..4628dd4 100644 --- a/tests/ollama-client-cache.test.ts +++ b/tests/ollama-client-cache.test.ts @@ -6,6 +6,7 @@ import { OllamaMessage, OllamaTool, CacheConfig } from '../src/types'; const mockInitialize = jest.fn().mockResolvedValue(undefined); const mockGetCache = jest.fn().mockResolvedValue(null); const mockSetCache = jest.fn().mockResolvedValue(undefined); +const mockClearCache = jest.fn().mockResolvedValue(undefined); // Mock the semantic cache service BEFORE importing OllamaClient jest.mock('../src/semantic-cache', () => ({ @@ -13,6 +14,7 @@ jest.mock('../src/semantic-cache', () => ({ initialize: mockInitialize, getCache: mockGetCache, setCache: mockSetCache, + clearCache: mockClearCache, })), })); @@ -68,6 +70,7 @@ describe('OllamaClient with Semantic Cache', () => { similarityThreshold: 0.85, collectionName: 'test_cache', embeddingModel: 'nomic-embed-text', + chromaUrl: 'http://localhost:8000', }; beforeEach(() => { @@ -101,6 +104,7 @@ describe('OllamaClient with Semantic Cache', () => { }); it('should not create cache service when no config provided', () => { + jest.clearAllMocks(); // Reset the call recorded by beforeEach before checking new OllamaClient('http://localhost:11434', 'llama3', mockFetch); expect(SemanticCacheService).not.toHaveBeenCalled(); @@ -350,19 +354,8 @@ describe('OllamaClient with Semantic Cache', () => { // Mock cache service to throw an error mockGetCache.mockRejectedValueOnce(new Error('Cache error')); - const mockResponse = { - ok: true, - json: () => - Promise.resolve({ - message: { - content: 'LLM response after cache failure', - }, - }), - }; - mockFetch.mockResolvedValueOnce(mockResponse); - - // The chat method does not handle cache errors, so it should propagate - // However, the client should still be usable + // Note: fetch is never reached because the cache throws first. + // The chat method does not handle cache errors, so it should propagate. await expect(client.chat(mockMessages)).rejects.toThrow('Cache error'); }); @@ -414,6 +407,18 @@ describe('OllamaClient with Semantic Cache', () => { expect(mockGetCache).toHaveBeenCalledWith('Second question'); }); + it('should call clearCache on the cache service', async () => { + await client.clearCache(); + expect(mockClearCache).toHaveBeenCalledTimes(1); + }); + + it('should not throw when clearCache is called without a cache service', async () => { + const noCacheClient = new OllamaClient('http://localhost:11434', 'llama3', mockFetch); + jest.clearAllMocks(); + await expect(noCacheClient.clearCache()).resolves.toBeUndefined(); + expect(mockClearCache).not.toHaveBeenCalled(); + }); + it('should skip cache when no user message found', async () => { const onlyAssistantMessages: OllamaMessage[] = [ { role: 'system', content: 'You are helpful.' }, diff --git a/tests/semantic-cache.test.ts b/tests/semantic-cache.test.ts index 4b00aa1..a09d017 100644 --- a/tests/semantic-cache.test.ts +++ b/tests/semantic-cache.test.ts @@ -7,12 +7,20 @@ import { CacheConfig } from '../src/types'; jest.mock('chromadb', () => ({ ChromaClient: jest.fn().mockImplementation(() => { return { - getOrCreateCollection: jest.fn().mockResolvedValue({ + getOrCreateCollection: jest.fn().mockReturnValue({ query: jest.fn(), add: jest.fn(), + upsert: jest.fn(), }), + deleteCollection: jest.fn(), }; }), + IncludeEnum: { + Documents: 'documents', + Embeddings: 'embeddings', + Metadatas: 'metadatas', + Distances: 'distances', + }, })); // Now import SemanticCacheService after mocking @@ -21,10 +29,12 @@ import { SemanticCacheService } from '../src/semantic-cache'; jest.spyOn(global, 'fetch').mockImplementation(jest.fn()); const mockChromaClient = { - getOrCreateCollection: jest.fn().mockResolvedValue({ + getOrCreateCollection: jest.fn().mockReturnValue({ query: jest.fn(), add: jest.fn(), + upsert: jest.fn(), }), + deleteCollection: jest.fn(), }; // Set up mock instance @@ -44,6 +54,7 @@ describe('SemanticCacheService', () => { similarityThreshold: 0.85, collectionName: 'test_cache', embeddingModel: 'nomic-embed-text', + chromaUrl: 'http://localhost:8000', }; service = new SemanticCacheService('http://localhost:11434', config); @@ -211,32 +222,30 @@ describe('SemanticCacheService', () => { json: () => Promise.resolve({ embedding: [0.1, 0.2, 0.3] }), }); - // crypto.randomUUID mock - const mockUuid = 'mock-uuid-123' as any; - jest.spyOn(crypto, 'randomUUID').mockReturnValue(mockUuid); - await service.setCache('test prompt', 'test response'); const mockCollection = mockChromaClient.getOrCreateCollection(); - expect(mockCollection.add).toHaveBeenCalledWith({ - ids: [mockUuid], - embeddings: [[0.1, 0.2, 0.3]], - metadatas: [{ fullResponse: 'test response' }], - }); + expect(mockCollection.upsert).toHaveBeenCalledWith( + expect.objectContaining({ + ids: [expect.any(String)], + embeddings: [[0.1, 0.2, 0.3]], + metadatas: [{ fullResponse: 'test response' }], + }) + ); }); it('should not add entry when prompt is empty', async () => { await service.setCache(' ', 'test response'); expect(mockFetch).not.toHaveBeenCalled(); - expect(mockChromaClient.getOrCreateCollection().add).not.toHaveBeenCalled(); + expect(mockChromaClient.getOrCreateCollection().upsert).not.toHaveBeenCalled(); }); it('should not add entry when response is empty', async () => { await service.setCache('test prompt', ' '); expect(mockFetch).not.toHaveBeenCalled(); - expect(mockChromaClient.getOrCreateCollection().add).not.toHaveBeenCalled(); + expect(mockChromaClient.getOrCreateCollection().upsert).not.toHaveBeenCalled(); }); it('should not add entry when cache is disabled', async () => { @@ -255,13 +264,33 @@ describe('SemanticCacheService', () => { json: () => Promise.resolve({ embedding: [0.1, 0.2, 0.3] }), }); - const mockUuid = 'mock-uuid-456' as any; - jest.spyOn(crypto, 'randomUUID').mockReturnValue(mockUuid); - - mockChromaClient.getOrCreateCollection().add.mockRejectedValueOnce(new Error('Add failed')); + mockChromaClient + .getOrCreateCollection() + .upsert.mockRejectedValueOnce(new Error('Add failed')); // Should not throw await expect(service.setCache('test prompt', 'test response')).resolves.toBeUndefined(); }); }); + + describe('clearCache', () => { + beforeEach(async () => { + await service.initialize(); + mockChromaClient.deleteCollection.mockResolvedValue(undefined); + }); + + it('should delete the collection and re-initialize', async () => { + await service.clearCache(); + + expect(mockChromaClient.deleteCollection).toHaveBeenCalledWith({ name: 'test_cache' }); + // getOrCreateCollection called once in beforeEach initialize, once in clearCache re-init + expect(mockChromaClient.getOrCreateCollection).toHaveBeenCalledTimes(2); + }); + + it('should propagate errors from deleteCollection', async () => { + mockChromaClient.deleteCollection.mockRejectedValueOnce(new Error('Delete failed')); + + await expect(service.clearCache()).rejects.toThrow('Delete failed'); + }); + }); });