{[
- { key: "actor-prose", label: "Actor Prose Generation", desc: "Generates roleplay/narrative prose for Non-Player Characters." },
- { key: "llm-validator", label: "LLM Validator", desc: "Arbitrates and validates proposed actions against the world state rules." },
- { key: "intent-decoder", label: "Intent Decoder", desc: "Splits raw prose actions into structured intents (Player and NPC)." },
- { key: "timedelta", label: "TimeDelta Generator", desc: "Calculates the duration of character actions to advance the game clock." },
+ { key: "actor-prose", label: "Actor Prose Generation", desc: "Generates roleplay/narrative prose for Non-Player Characters.", type: "generative" },
+ { key: "llm-validator", label: "LLM Validator", desc: "Arbitrates and validates proposed actions against the world state rules.", type: "generative" },
+ { key: "intent-decoder", label: "Intent Decoder", desc: "Splits raw prose actions into structured intents (Player and NPC).", type: "generative" },
+ { key: "timedelta", label: "TimeDelta Generator", desc: "Calculates the duration of character actions to advance the game clock.", type: "generative" },
+ { key: "embeddings", label: "Text Embeddings Generator", desc: "Generates vector embeddings for long-term memory retrieval.", type: "embedding" },
].map((task) => (
- {instances.map((inst) => (
-
- ))}
+ {instances
+ .filter((inst) => (inst.type || "generative") === task.type)
+ .map((inst) => (
+
+ ))}
))}
diff --git a/apps/gui/src/app/play/actions.ts b/apps/gui/src/app/play/actions.ts
index c36a26e..a04b4f5 100644
--- a/apps/gui/src/app/play/actions.ts
+++ b/apps/gui/src/app/play/actions.ts
@@ -4,7 +4,7 @@ import path from "path";
import fs from "fs";
import { simulationManager } from "@/lib/simulation";
import type { SimSnapshot } from "@/lib/simulation";
-import { ProviderManager, LLMProviderInstance, AVAILABLE_PROVIDERS, LLMProviderMeta } from "@omnia/llm";
+import { ProviderManager, ModelProviderInstance, AVAILABLE_PROVIDERS, ModelProviderMeta } from "@omnia/llm";
function resolveScenarioPath(relative: string): string {
const cwd = process.cwd();
@@ -233,7 +233,7 @@ export async function deleteSimulation(simId: string): Promise<
}
}
-export async function listProviderInstances(): Promise
{
+export async function listProviderInstances(): Promise {
return ProviderManager.list();
}
@@ -242,8 +242,9 @@ export async function createProviderInstance(
providerName: string,
apiKey: string,
modelName?: string,
-): Promise {
- return ProviderManager.create(name, providerName, apiKey, modelName);
+ type: "generative" | "embedding" = "generative",
+): Promise {
+ return ProviderManager.create(name, providerName, apiKey, modelName, type);
}
export async function deleteProviderInstance(id: string): Promise {
@@ -260,8 +261,9 @@ export async function updateProviderInstance(
providerName: string,
apiKey?: string,
modelName?: string,
+ type: "generative" | "embedding" = "generative",
): Promise {
- ProviderManager.update(id, name, providerName, apiKey, modelName);
+ ProviderManager.update(id, name, providerName, apiKey, modelName, type);
}
export async function getProviderMappings(): Promise> {
@@ -275,6 +277,10 @@ export async function setProviderMapping(
ProviderManager.setMapping(task, providerInstanceId);
}
-export async function getAvailableProviders(): Promise {
+export async function getAvailableProviders(): Promise {
return AVAILABLE_PROVIDERS;
}
+
+export async function regenerateEmbeddings(newProviderInstanceId?: string): Promise {
+ await simulationManager.regenerateAllEmbeddings(newProviderInstanceId);
+}
diff --git a/apps/gui/src/lib/simulation.ts b/apps/gui/src/lib/simulation.ts
index d824039..990b073 100644
--- a/apps/gui/src/lib/simulation.ts
+++ b/apps/gui/src/lib/simulation.ts
@@ -25,7 +25,7 @@ import {
IActorProseGenerator,
buildBufferEntryForIntent,
} from "@omnia/actor";
-import { GeminiProvider, ILLMProvider, MockLLMProvider, ProviderManager, OpenRouterProvider } from "@omnia/llm";
+import { GeminiProvider, ILLMProvider, MockLLMProvider, ProviderManager, OpenRouterProvider, IEmbeddingProvider, GeminiEmbeddingProvider, MockEmbeddingProvider, ModelProviderInstance } from "@omnia/llm";
import { ScenarioLoader } from "@omnia/scenario";
import type {
@@ -102,6 +102,7 @@ interface SimSession {
validatorProvider: ILLMProvider;
decoderProvider: ILLMProvider;
timedeltaProvider: ILLMProvider;
+ embeddingProvider: IEmbeddingProvider;
architect: Architect;
aliasGenerator: AliasDeltaGenerator;
log: LogEntry[];
@@ -120,14 +121,14 @@ class SimulationManager {
playEntityName?: string,
providerInstanceId?: string,
): Promise {
- let activeInstance = providerInstanceId
- ? ProviderManager.list().find((p) => p.id === providerInstanceId)
- : ProviderManager.getActive();
+ let activeInstance: ModelProviderInstance | null = providerInstanceId
+ ? ProviderManager.list().find((p) => p.id === providerInstanceId) || null
+ : ProviderManager.getActive("generative");
if (!activeInstance) {
const envKey = process.env.GOOGLE_API_KEY;
if (envKey) {
- activeInstance = ProviderManager.create("Default (Env)", "google-genai", envKey);
+ activeInstance = ProviderManager.create("Default (Env)", "google-genai", envKey, undefined, "generative");
}
}
@@ -221,13 +222,13 @@ class SimulationManager {
}
const list = ProviderManager.list();
- const active = ProviderManager.getActive() || activeInstance;
+ const active = ProviderManager.getActive("generative") || activeInstance;
const mappings = ProviderManager.getMappings();
const resolveProviderForTask = (task: string): ILLMProvider => {
const mappedId = mappings[task];
let inst = mappedId ? list.find((p) => p.id === mappedId) : null;
- if (!inst) {
+ if (!inst || inst.type !== "generative") {
inst = active;
}
@@ -244,10 +245,29 @@ class SimulationManager {
}
};
+ const resolveEmbeddingProvider = (): IEmbeddingProvider => {
+ const mappedId = mappings["embeddings"];
+ let inst = mappedId ? list.find((p) => p.id === mappedId) : null;
+ if (!inst || inst.type !== "embedding") {
+ inst = ProviderManager.getActive("embedding");
+ }
+
+ const key = inst ? inst.apiKey : (process.env.GOOGLE_API_KEY || "");
+ const providerName = inst ? inst.providerName : "google-genai";
+ const modelName = inst ? inst.modelName : undefined;
+
+ if (providerName === "google-genai") {
+ return new GeminiEmbeddingProvider(key, modelName);
+ } else {
+ return new MockEmbeddingProvider(modelName);
+ }
+ };
+
const actorProvider = resolveProviderForTask("actor-prose");
const validatorProvider = resolveProviderForTask("llm-validator");
const decoderProvider = resolveProviderForTask("intent-decoder");
const timedeltaProvider = resolveProviderForTask("timedelta");
+ const embeddingProvider = resolveEmbeddingProvider();
const architect = new Architect(
{ validator: validatorProvider, timedelta: timedeltaProvider },
@@ -261,9 +281,9 @@ class SimulationManager {
coreRepo,
bufferRepo,
ledgerRepo,
- worldInstanceId,
+ worldInstanceId: worldInstanceId,
scenarioName: scenarioJson.name,
- scenarioDescription: scenarioJson.description,
+ scenarioDescription: scenarioJson.description || "",
turn: 1,
maxTurns: 20,
entities: entityInfos,
@@ -273,6 +293,7 @@ class SimulationManager {
validatorProvider,
decoderProvider,
timedeltaProvider,
+ embeddingProvider,
architect,
aliasGenerator,
log: [],
@@ -655,19 +676,19 @@ class SimulationManager {
}
const list = ProviderManager.list();
- const active = ProviderManager.getActive();
+ const active = ProviderManager.getActive("generative");
const mappings = state.providerMappings || {};
const resolveProviderForTask = (task: string): ILLMProvider => {
const mappedId = mappings[task];
let inst = mappedId ? list.find((p) => p.id === mappedId) : null;
- if (!inst) {
+ if (!inst || inst.type !== "generative") {
inst = active;
}
if (!inst) {
const envKey = process.env.GOOGLE_API_KEY;
if (envKey) {
- inst = ProviderManager.create("Default (Env)", "google-genai", envKey);
+ inst = ProviderManager.create("Default (Env)", "google-genai", envKey, undefined, "generative");
}
}
@@ -684,6 +705,30 @@ class SimulationManager {
}
};
+ const resolveEmbeddingProvider = (): IEmbeddingProvider => {
+ const mappedId = mappings["embeddings"];
+ let inst = mappedId ? list.find((p) => p.id === mappedId) : null;
+ if (!inst || inst.type !== "embedding") {
+ inst = ProviderManager.getActive("embedding");
+ }
+ if (!inst) {
+ const envKey = process.env.GOOGLE_API_KEY;
+ if (envKey) {
+ inst = ProviderManager.create("Default Embed (Env)", "google-genai", envKey, "gemini-embedding-001", "embedding");
+ }
+ }
+
+ if (!inst) {
+ throw new Error(`No active Embedding Provider Instance found for task "embeddings". Please configure an embedding key in Settings first.`);
+ }
+
+ if (inst.providerName === "google-genai") {
+ return new GeminiEmbeddingProvider(inst.apiKey, inst.modelName);
+ } else {
+ return new MockEmbeddingProvider(inst.modelName);
+ }
+ };
+
const coreRepo = new SQLiteRepository(db);
const bufferRepo = new BufferRepository(db);
const ledgerRepo = new LedgerRepository(db);
@@ -692,6 +737,7 @@ class SimulationManager {
const validatorProvider = resolveProviderForTask("llm-validator");
const decoderProvider = resolveProviderForTask("intent-decoder");
const timedeltaProvider = resolveProviderForTask("timedelta");
+ const embeddingProvider = resolveEmbeddingProvider();
const architect = new Architect(
{ validator: validatorProvider, timedelta: timedeltaProvider },
@@ -717,6 +763,7 @@ class SimulationManager {
validatorProvider,
decoderProvider,
timedeltaProvider,
+ embeddingProvider,
architect,
aliasGenerator,
log: state.log || [],
@@ -784,6 +831,53 @@ class SimulationManager {
});
}
+ async regenerateAllEmbeddings(newProviderInstanceId?: string): Promise {
+ const dbDir = path.resolve(process.cwd(), "data");
+ if (!fs.existsSync(dbDir)) return;
+
+ const files = fs.readdirSync(dbDir).filter(f => f.startsWith("sim-") && f.endsWith(".db"));
+
+ const list = ProviderManager.list();
+ let inst = newProviderInstanceId ? list.find((p) => p.id === newProviderInstanceId) : null;
+ if (!inst || inst.type !== "embedding") {
+ inst = ProviderManager.getActive("embedding");
+ }
+
+ const key = inst ? inst.apiKey : (process.env.GOOGLE_API_KEY || "");
+ const providerName = inst ? inst.providerName : "google-genai";
+ const modelName = inst ? inst.modelName : undefined;
+
+ let embeddingProvider: IEmbeddingProvider;
+ if (providerName === "google-genai") {
+ embeddingProvider = new GeminiEmbeddingProvider(key, modelName);
+ } else {
+ embeddingProvider = new MockEmbeddingProvider(modelName);
+ }
+
+ for (const file of files) {
+ const dbPath = path.join(dbDir, file);
+ const id = file.replace(".db", "");
+ const activeSession = this.sessions.get(id);
+ const db = activeSession ? activeSession.db : new Database(dbPath);
+
+ try {
+ const rows = db.prepare(`SELECT id, content FROM ledger_entries`).all() as { id: string; content: string }[];
+
+ for (const row of rows) {
+ const vector = await embeddingProvider.embed(row.content);
+ const buffer = Buffer.from(new Float32Array(vector).buffer);
+ db.prepare(`UPDATE ledger_entries SET embedding = ? WHERE id = ?`).run(buffer, row.id);
+ }
+ } catch (err) {
+ console.error(`Failed to regenerate embeddings for ${file}:`, err);
+ } finally {
+ if (!activeSession) {
+ db.close();
+ }
+ }
+ }
+ }
+
private save(session: SimSession): void {
const state: SavedState = {
scenarioName: session.scenarioName,
diff --git a/packages/llm/src/llm.ts b/packages/llm/src/llm.ts
index 5a1d3e8..4658292 100644
--- a/packages/llm/src/llm.ts
+++ b/packages/llm/src/llm.ts
@@ -30,46 +30,55 @@ export interface LLMCallRecord {
export interface ILLMProvider {
providerName: string;
- // We use Zod to ensure the generic T matches the schema
generateStructuredResponse(
request: LLMRequest,
): Promise>>;
lastCalls?: LLMCallRecord[];
}
-export interface LLMProviderInstance {
+export interface IEmbeddingProvider {
+ providerName: string;
+ embed(text: string): Promise;
+}
+
+export interface ModelProviderInstance {
id: string;
name: string;
providerName: string;
apiKey: string;
isActive: boolean;
modelName?: string;
+ type: "generative" | "embedding";
}
-export interface LLMProviderMeta {
+export interface ModelProviderMeta {
id: string;
displayName: string;
description: string;
defaultModel: string;
+ defaultEmbeddingModel: string;
}
-export const AVAILABLE_PROVIDERS: LLMProviderMeta[] = [
+export const AVAILABLE_PROVIDERS: ModelProviderMeta[] = [
{
id: "google-genai",
displayName: "Google Gemini",
description: "Official Gemini integration using Google Gen AI SDK",
defaultModel: "gemini-2.5-flash",
+ defaultEmbeddingModel: "gemini-embedding-001",
},
{
id: "openrouter",
displayName: "OpenRouter",
description: "Multi-model router supporting Anthropic, OpenAI, DeepSeek, and local models",
defaultModel: "google/gemini-2.5-flash",
+ defaultEmbeddingModel: "openai/text-embedding-3-small",
},
{
id: "mock",
displayName: "Mock LLM Provider",
description: "Stateless mock provider for testing and offline development",
defaultModel: "mock",
+ defaultEmbeddingModel: "mock-embeddings",
},
];
diff --git a/packages/llm/src/provider-manager.ts b/packages/llm/src/provider-manager.ts
index e032aaa..f953935 100644
--- a/packages/llm/src/provider-manager.ts
+++ b/packages/llm/src/provider-manager.ts
@@ -1,7 +1,7 @@
import Database from "better-sqlite3";
import path from "path";
import fs from "fs";
-import type { LLMProviderInstance } from "./llm.js";
+import type { ModelProviderInstance } from "./llm.js";
function getWorkspaceRoot() {
let current = process.cwd();
@@ -35,7 +35,8 @@ function getSettingsDb() {
providerName TEXT NOT NULL,
apiKey TEXT NOT NULL,
isActive INTEGER NOT NULL DEFAULT 0,
- modelName TEXT
+ modelName TEXT,
+ type TEXT NOT NULL DEFAULT 'generative'
)
`).run();
@@ -44,12 +45,18 @@ function getSettingsDb() {
} catch {
// ignore
}
+
+ try {
+ db.prepare(`ALTER TABLE provider_instances ADD COLUMN type TEXT NOT NULL DEFAULT 'generative'`).run();
+ } catch {
+ // ignore
+ }
return db;
}
export class ProviderManager {
- static list(): LLMProviderInstance[] {
+ static list(): ModelProviderInstance[] {
const db = getSettingsDb();
try {
const rows = db.prepare(`SELECT * FROM provider_instances`).all() as {
@@ -59,6 +66,7 @@ export class ProviderManager {
apiKey: string;
isActive: number;
modelName?: string;
+ type: string;
}[];
return rows.map((r) => ({
id: r.id,
@@ -67,25 +75,34 @@ export class ProviderManager {
apiKey: r.apiKey,
isActive: r.isActive === 1,
modelName: r.modelName || undefined,
+ type: (r.type as "generative" | "embedding") || "generative",
}));
} finally {
db.close();
}
}
- static create(name: string, providerName: string, apiKey: string, modelName?: string): LLMProviderInstance {
+ static create(
+ name: string,
+ providerName: string,
+ apiKey: string,
+ modelName?: string,
+ type: "generative" | "embedding" = "generative"
+ ): ModelProviderInstance {
const db = getSettingsDb();
try {
const id = "provider-" + Date.now();
- const activeCount = db.prepare(`SELECT COUNT(*) as count FROM provider_instances WHERE isActive = 1`).get() as { count: number };
+ const activeCount = db
+ .prepare(`SELECT COUNT(*) as count FROM provider_instances WHERE isActive = 1 AND type = ?`)
+ .get(type) as { count: number };
const isActive = activeCount.count === 0 ? 1 : 0;
db.prepare(`
- INSERT INTO provider_instances (id, name, providerName, apiKey, isActive, modelName)
- VALUES (?, ?, ?, ?, ?, ?)
- `).run(id, name, providerName, apiKey, isActive, modelName || null);
+ INSERT INTO provider_instances (id, name, providerName, apiKey, isActive, modelName, type)
+ VALUES (?, ?, ?, ?, ?, ?, ?)
+ `).run(id, name, providerName, apiKey, isActive, modelName || null, type);
- return { id, name, providerName, apiKey, isActive: isActive === 1, modelName };
+ return { id, name, providerName, apiKey, isActive: isActive === 1, modelName, type };
} finally {
db.close();
}
@@ -94,11 +111,13 @@ export class ProviderManager {
static delete(id: string): void {
const db = getSettingsDb();
try {
- const provider = db.prepare(`SELECT isActive FROM provider_instances WHERE id = ?`).get(id) as { isActive: number } | undefined;
+ const provider = db.prepare(`SELECT isActive, type FROM provider_instances WHERE id = ?`).get(id) as { isActive: number; type: string } | undefined;
db.prepare(`DELETE FROM provider_instances WHERE id = ?`).run(id);
if (provider && provider.isActive === 1) {
- const next = db.prepare(`SELECT id FROM provider_instances LIMIT 1`).get() as { id: string } | undefined;
+ const next = db
+ .prepare(`SELECT id FROM provider_instances WHERE type = ? LIMIT 1`)
+ .get(provider.type) as { id: string } | undefined;
if (next) {
db.prepare(`UPDATE provider_instances SET isActive = 1 WHERE id = ?`).run(next.id);
}
@@ -111,60 +130,85 @@ export class ProviderManager {
static setActive(id: string): void {
const db = getSettingsDb();
try {
- db.prepare(`UPDATE provider_instances SET isActive = 0`).run();
- db.prepare(`UPDATE provider_instances SET isActive = 1 WHERE id = ?`).run(id);
- } finally {
- db.close();
- }
- }
-
- static update(id: string, name: string, providerName: string, apiKey?: string, modelName?: string): void {
- const db = getSettingsDb();
- try {
- if (apiKey && apiKey.trim()) {
- db.prepare(`
- UPDATE provider_instances
- SET name = ?, providerName = ?, apiKey = ?, modelName = ?
- WHERE id = ?
- `).run(name, providerName, apiKey, modelName || null, id);
- } else {
- db.prepare(`
- UPDATE provider_instances
- SET name = ?, providerName = ?, modelName = ?
- WHERE id = ?
- `).run(name, providerName, modelName || null, id);
+ const target = db.prepare(`SELECT type FROM provider_instances WHERE id = ?`).get(id) as { type: string } | undefined;
+ if (target) {
+ db.prepare(`UPDATE provider_instances SET isActive = 0 WHERE type = ?`).run(target.type);
+ db.prepare(`UPDATE provider_instances SET isActive = 1 WHERE id = ?`).run(id);
}
} finally {
db.close();
}
}
- static getActive(): LLMProviderInstance | null {
+ static update(
+ id: string,
+ name: string,
+ providerName: string,
+ apiKey?: string,
+ modelName?: string,
+ type: "generative" | "embedding" = "generative"
+ ): void {
const db = getSettingsDb();
try {
- // Query the DB
- const row = db.prepare(`SELECT * FROM provider_instances WHERE isActive = 1`).get() as {
+ if (apiKey && apiKey.trim()) {
+ db.prepare(`
+ UPDATE provider_instances
+ SET name = ?, providerName = ?, apiKey = ?, modelName = ?, type = ?
+ WHERE id = ?
+ `).run(name, providerName, apiKey, modelName || null, type, id);
+ } else {
+ db.prepare(`
+ UPDATE provider_instances
+ SET name = ?, providerName = ?, modelName = ?, type = ?
+ WHERE id = ?
+ `).run(name, providerName, modelName || null, type, id);
+ }
+ } finally {
+ db.close();
+ }
+ }
+
+ static getActive(type: "generative" | "embedding" = "generative"): ModelProviderInstance | null {
+ const db = getSettingsDb();
+ try {
+ const row = db.prepare(`SELECT * FROM provider_instances WHERE isActive = 1 AND type = ?`).get(type) as {
id: string;
name: string;
providerName: string;
apiKey: string;
isActive: number;
modelName?: string;
+ type: string;
} | undefined;
if (!row) {
- // Check if there are any rows at all
const totalCount = db.prepare(`SELECT COUNT(*) as count FROM provider_instances`).get() as { count: number };
if (totalCount.count === 0) {
- // Database is completely empty! Check if GOOGLE_API_KEY env is set.
const envKey = process.env.GOOGLE_API_KEY;
if (envKey && envKey.trim()) {
- // Auto-bootstrap default active instance from env
const id = "provider-default-env";
db.prepare(`
- INSERT INTO provider_instances (id, name, providerName, apiKey, isActive, modelName)
- VALUES (?, ?, ?, ?, ?, ?)
- `).run(id, "Default (Env)", "google-genai", envKey, 1, "gemini-2.5-flash");
+ INSERT INTO provider_instances (id, name, providerName, apiKey, isActive, modelName, type)
+ VALUES (?, ?, ?, ?, ?, ?, ?)
+ `).run(id, "Default (Env)", "google-genai", envKey, 1, "gemini-2.5-flash", "generative");
+
+ const embedId = "provider-default-env-embed";
+ db.prepare(`
+ INSERT INTO provider_instances (id, name, providerName, apiKey, isActive, modelName, type)
+ VALUES (?, ?, ?, ?, ?, ?, ?)
+ `).run(embedId, "Default Embed (Env)", "google-genai", envKey, 1, "gemini-embedding-001", "embedding");
+
+ if (type === "embedding") {
+ return {
+ id: embedId,
+ name: "Default Embed (Env)",
+ providerName: "google-genai",
+ apiKey: envKey,
+ isActive: true,
+ modelName: "gemini-embedding-001",
+ type: "embedding",
+ };
+ }
return {
id,
@@ -173,6 +217,7 @@ export class ProviderManager {
apiKey: envKey,
isActive: true,
modelName: "gemini-2.5-flash",
+ type: "generative",
};
}
}
@@ -186,11 +231,22 @@ export class ProviderManager {
apiKey: row.apiKey,
isActive: true,
modelName: row.modelName || undefined,
+ type: (row.type as "generative" | "embedding") || "generative",
};
} catch {
- // Lock or write issue fallback: return an in-memory active key if env key exists
const envKey = process.env.GOOGLE_API_KEY;
if (envKey) {
+ if (type === "embedding") {
+ return {
+ id: "provider-default-env-embed-fallback",
+ name: "Default Embed (Env Fallback)",
+ providerName: "google-genai",
+ apiKey: envKey,
+ isActive: true,
+ modelName: "gemini-embedding-001",
+ type: "embedding",
+ };
+ }
return {
id: "provider-default-env-fallback",
name: "Default (Env Fallback)",
@@ -198,6 +254,7 @@ export class ProviderManager {
apiKey: envKey,
isActive: true,
modelName: "gemini-2.5-flash",
+ type: "generative",
};
}
return null;
diff --git a/packages/llm/src/providers/google-genai.ts b/packages/llm/src/providers/google-genai.ts
index 17b0d1f..caf3c14 100644
--- a/packages/llm/src/providers/google-genai.ts
+++ b/packages/llm/src/providers/google-genai.ts
@@ -1,6 +1,6 @@
import { z } from "zod";
-import { ChatGoogleGenerativeAI } from "@langchain/google-genai";
-import { ILLMProvider, LLMRequest, LLMResponse, LLMCallRecord } from "../llm.js";
+import { ChatGoogleGenerativeAI, GoogleGenerativeAIEmbeddings } from "@langchain/google-genai";
+import { ILLMProvider, LLMRequest, LLMResponse, LLMCallRecord, IEmbeddingProvider } from "../llm.js";
import { llmConfig } from "../config.js";
import { ProviderManager } from "../provider-manager.js";
@@ -19,7 +19,7 @@ export class GeminiProvider implements ILLMProvider {
let model = modelName;
if (!key) {
- const active = ProviderManager.getActive();
+ const active = ProviderManager.getActive("generative");
if (active) {
key = active.apiKey;
if (!model) {
@@ -78,3 +78,43 @@ export class GeminiProvider implements ILLMProvider {
return { success: true, data: parsed, usage };
}
}
+
+export class GeminiEmbeddingProvider implements IEmbeddingProvider {
+ static readonly providerId = "google-genai";
+ static readonly displayName = "Google Gemini Embeddings";
+
+ providerName = "Gemini";
+ private model: GoogleGenerativeAIEmbeddings;
+
+ constructor(apiKey?: string, modelName?: string) {
+ let key = apiKey;
+ let model = modelName;
+
+ if (!key) {
+ const active = ProviderManager.getActive("embedding");
+ if (active) {
+ key = active.apiKey;
+ if (!model) {
+ model = active.modelName;
+ }
+ }
+ }
+
+ if (!key) {
+ key = llmConfig.GOOGLE_API_KEY;
+ }
+
+ if (!key) {
+ throw new Error("GOOGLE_API_KEY is required to initialize GeminiEmbeddingProvider");
+ }
+
+ this.model = new GoogleGenerativeAIEmbeddings({
+ apiKey: key,
+ modelName: model || "gemini-embedding-001",
+ });
+ }
+
+ async embed(text: string): Promise {
+ return this.model.embedQuery(text);
+ }
+}
diff --git a/packages/llm/src/providers/mock.ts b/packages/llm/src/providers/mock.ts
index f11a90d..53b2fbc 100644
--- a/packages/llm/src/providers/mock.ts
+++ b/packages/llm/src/providers/mock.ts
@@ -1,5 +1,5 @@
import { z } from "zod";
-import { ILLMProvider, LLMRequest, LLMResponse, LLMCallRecord } from "../llm.js";
+import { ILLMProvider, LLMRequest, LLMResponse, LLMCallRecord, IEmbeddingProvider } from "../llm.js";
export class MockLLMProvider implements ILLMProvider {
static readonly providerId = "mock";
@@ -34,3 +34,21 @@ export class MockLLMProvider implements ILLMProvider {
}
}
}
+
+export class MockEmbeddingProvider implements IEmbeddingProvider {
+ static readonly providerId = "mock";
+
+ providerName = "mock";
+
+ constructor(private modelName?: string) {}
+
+ async embed(text: string): Promise {
+ // Return a deterministic mock 768-dimensional vector based on the text
+ const vec = new Array(768).fill(0).map((_, i) => {
+ // Return a predictable float between -1.0 and 1.0
+ const charCode = text.charCodeAt(i % text.length) || 0;
+ return Math.sin(charCode + i);
+ });
+ return vec;
+ }
+}
diff --git a/packages/llm/src/providers/openrouter.ts b/packages/llm/src/providers/openrouter.ts
index 200ac4e..e8a2bdd 100644
--- a/packages/llm/src/providers/openrouter.ts
+++ b/packages/llm/src/providers/openrouter.ts
@@ -19,7 +19,7 @@ export class OpenRouterProvider implements ILLMProvider {
let model = modelName;
if (!key) {
- const active = ProviderManager.getActive();
+ const active = ProviderManager.getActive("generative");
if (active) {
key = active.apiKey;
if (!model) {
diff --git a/packages/llm/tests/mock.test.ts b/packages/llm/tests/mock.test.ts
index d9455c4..9b065f4 100644
--- a/packages/llm/tests/mock.test.ts
+++ b/packages/llm/tests/mock.test.ts
@@ -1,6 +1,6 @@
import { describe, test, expect } from "vitest";
import { z } from "zod";
-import { MockLLMProvider } from "@omnia/llm";
+import { MockLLMProvider, MockEmbeddingProvider } from "@omnia/llm";
describe("MockLLMProvider Unit Tests (Tier 1)", () => {
test("returns parsed matching data for valid mock response", async () => {
@@ -61,3 +61,21 @@ describe("MockLLMProvider Unit Tests (Tier 1)", () => {
expect(response.data).toBeUndefined();
});
});
+
+describe("MockEmbeddingProvider Unit Tests (Tier 1)", () => {
+ test("generates deterministic 768-dimensional vectors", async () => {
+ const provider = new MockEmbeddingProvider("mock-embeddings");
+ const text = "Hello world";
+ const vec1 = await provider.embed(text);
+ const vec2 = await provider.embed(text);
+
+ expect(vec1.length).toBe(768);
+ expect(vec2.length).toBe(768);
+ expect(vec1).toEqual(vec2); // Deterministic
+
+ // Ensure values are numbers between -1.0 and 1.0 (since they are generated with Math.sin)
+ expect(typeof vec1[0]).toBe("number");
+ expect(vec1[0]).toBeGreaterThanOrEqual(-1.0);
+ expect(vec1[0]).toBeLessThanOrEqual(1.0);
+ });
+});