major(llm): Support Embedding Providers

This commit is contained in:
2026-07-11 14:31:18 +05:30
parent a5fa43e2e6
commit 597d4f1711
10 changed files with 398 additions and 89 deletions

View File

@@ -1,6 +1,6 @@
/// <reference types="next" />
/// <reference types="next/image-types/global" />
import "./.next/types/routes.d.ts";
import "./.next/dev/types/routes.d.ts";
// NOTE: This file should not be edited
// see https://nextjs.org/docs/app/api-reference/config/typescript for more information.

View File

@@ -11,8 +11,9 @@ import {
setProviderMapping,
updateProviderInstance,
getAvailableProviders,
regenerateEmbeddings,
} from "@/app/play/actions";
import type { LLMProviderInstance, LLMProviderMeta } from "@omnia/llm";
import type { ModelProviderInstance, ModelProviderMeta } from "@omnia/llm";
interface ConfigStatus {
apiKeySet: boolean;
@@ -23,9 +24,9 @@ interface ConfigStatus {
export default function ConfigPage() {
const [config, setConfig] = useState<ConfigStatus | null>(null);
const [instances, setInstances] = useState<LLMProviderInstance[]>([]);
const [instances, setInstances] = useState<ModelProviderInstance[]>([]);
const [mappings, setMappings] = useState<Record<string, string>>({});
const [availableProviders, setAvailableProviders] = useState<LLMProviderMeta[]>([]);
const [availableProviders, setAvailableProviders] = useState<ModelProviderMeta[]>([]);
const [loading, setLoading] = useState(true);
const [error, setError] = useState("");
@@ -35,6 +36,7 @@ export default function ConfigPage() {
const [editKey, setEditKey] = useState("");
const [editModel, setEditModel] = useState("gemini-2.5-flash");
const [editIsActive, setEditIsActive] = useState(false);
const [editType, setEditType] = useState<"generative" | "embedding">("generative");
useEffect(() => {
if (selectedInstanceId === "new") {
@@ -42,6 +44,7 @@ export default function ConfigPage() {
const defaultProvider = "google-genai";
setEditProvider(defaultProvider);
setEditKey("");
setEditType("generative");
const pMeta = availableProviders.find((p) => p.id === defaultProvider);
setEditModel(pMeta?.defaultModel || "gemini-2.5-flash");
setEditIsActive(false);
@@ -51,8 +54,9 @@ export default function ConfigPage() {
setEditName(inst.name);
setEditProvider(inst.providerName);
setEditKey("");
setEditType(inst.type || "generative");
const pMeta = availableProviders.find((p) => p.id === inst.providerName);
setEditModel(inst.modelName || pMeta?.defaultModel || "gemini-2.5-flash");
setEditModel(inst.modelName || (inst.type === "embedding" ? pMeta?.defaultEmbeddingModel : pMeta?.defaultModel) || "gemini-2.5-flash");
setEditIsActive(inst.isActive);
}
}
@@ -62,7 +66,15 @@ export default function ConfigPage() {
setEditProvider(providerId);
const pMeta = availableProviders.find((p) => p.id === providerId);
if (pMeta) {
setEditModel(pMeta.defaultModel);
setEditModel(editType === "embedding" ? pMeta.defaultEmbeddingModel : pMeta.defaultModel);
}
};
const handleTypeChange = (type: "generative" | "embedding") => {
setEditType(type);
const pMeta = availableProviders.find((p) => p.id === editProvider);
if (pMeta) {
setEditModel(type === "embedding" ? pMeta.defaultEmbeddingModel : pMeta.defaultModel);
}
};
@@ -116,19 +128,42 @@ export default function ConfigPage() {
setLoading(true);
setError("");
let shouldRegenerate = false;
let targetInstanceId = selectedInstanceId;
if (selectedInstanceId === "new") {
if (!editKey.trim()) {
setError("API Key is required for new instances.");
setLoading(false);
return;
}
const created = await createProviderInstance(editName, editProvider, editKey, editModel || undefined);
const created = await createProviderInstance(editName, editProvider, editKey, editModel || undefined, editType);
if (editIsActive) {
await setActiveProviderInstance(created.id);
}
targetInstanceId = created.id;
setSelectedInstanceId(created.id);
} else {
await updateProviderInstance(selectedInstanceId, editName, editProvider, editKey || undefined, editModel || undefined);
const inst = instances.find((i) => i.id === selectedInstanceId);
if (inst && inst.type === "embedding") {
const isMapped = mappings["embeddings"] === selectedInstanceId;
const isActive = inst.isActive && !mappings["embeddings"];
if (isMapped || isActive) {
const hasChanged = inst.providerName !== editProvider || inst.modelName !== editModel;
if (hasChanged) {
const confirmChange = window.confirm(
"You have changed the configuration of the active embedding provider. This will delete all existing embeddings and regenerate them from scratch. Are you sure you want to do this?"
);
if (!confirmChange) {
setLoading(false);
return;
}
shouldRegenerate = true;
}
}
}
await updateProviderInstance(selectedInstanceId, editName, editProvider, editKey || undefined, editModel || undefined, editType);
if (editIsActive) {
await setActiveProviderInstance(selectedInstanceId);
}
@@ -136,6 +171,10 @@ export default function ConfigPage() {
await loadInstances();
await loadMappings();
if (shouldRegenerate && targetInstanceId !== "new") {
await regenerateEmbeddings(targetInstanceId);
}
} catch (err) {
setError(err instanceof Error ? err.message : String(err));
} finally {
@@ -162,9 +201,19 @@ export default function ConfigPage() {
};
const handleUpdateMapping = async (task: string, providerInstanceId: string) => {
if (task === "embeddings" && mappings[task] !== providerInstanceId) {
const confirmChange = window.confirm(
"Changing the embeddings provider will delete all existing embeddings and regenerate them from scratch. Are you sure you want to do this?"
);
if (!confirmChange) return;
}
try {
setLoading(true);
await setProviderMapping(task, providerInstanceId);
if (task === "embeddings") {
await regenerateEmbeddings(providerInstanceId);
}
await loadMappings();
} catch (err) {
setError(err instanceof Error ? err.message : String(err));
@@ -219,7 +268,7 @@ export default function ConfigPage() {
>
<div className="text-sm font-medium text-[#111]">{inst.name}</div>
<div className="mt-1 flex items-center justify-between text-xs text-gray-500">
<span>{inst.providerName}</span>
<span>{inst.providerName} ({inst.type || "generative"})</span>
{inst.isActive && (
<span className="rounded-full bg-green-100 px-1.5 py-[1px] text-[0.65rem] font-semibold text-green-700">
Active
@@ -257,6 +306,21 @@ export default function ConfigPage() {
/>
</div>
<div className="flex flex-col gap-1.5">
<label htmlFor="formType" className="text-xs font-medium text-gray-700">
Instance Type
</label>
<select
id="formType"
value={editType}
onChange={(e) => handleTypeChange(e.target.value as "generative" | "embedding")}
className="w-full rounded-md border border-gray-300 bg-white px-3 py-2 text-sm outline-none transition-[border-color,box-shadow] focus:border-blue-500 focus:ring-3 focus:ring-blue-500/15"
>
<option value="generative">Generative (Chat / Text Completion)</option>
<option value="embedding">Embedding (Vector generation)</option>
</select>
</div>
<div className="flex flex-col gap-1.5">
<label htmlFor="formProvider" className="text-xs font-medium text-gray-700">
Provider Type
@@ -364,10 +428,11 @@ export default function ConfigPage() {
</p>
<div className="mt-4 grid grid-cols-1 gap-4 md:grid-cols-2">
{[
{ 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) => (
<div
key={task.key}
@@ -383,11 +448,13 @@ export default function ConfigPage() {
className="w-full rounded border border-gray-300 bg-white px-2 py-1.5 text-xs"
>
<option value="">-- Use Active Key (Default) --</option>
{instances.map((inst) => (
<option key={inst.id} value={inst.id}>
{inst.name} ({inst.providerName}){inst.isActive ? " [Active]" : ""}
</option>
))}
{instances
.filter((inst) => (inst.type || "generative") === task.type)
.map((inst) => (
<option key={inst.id} value={inst.id}>
{inst.name} ({inst.providerName}){inst.isActive ? " [Active]" : ""}
</option>
))}
</select>
</div>
))}

View File

@@ -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<LLMProviderInstance[]> {
export async function listProviderInstances(): Promise<ModelProviderInstance[]> {
return ProviderManager.list();
}
@@ -242,8 +242,9 @@ export async function createProviderInstance(
providerName: string,
apiKey: string,
modelName?: string,
): Promise<LLMProviderInstance> {
return ProviderManager.create(name, providerName, apiKey, modelName);
type: "generative" | "embedding" = "generative",
): Promise<ModelProviderInstance> {
return ProviderManager.create(name, providerName, apiKey, modelName, type);
}
export async function deleteProviderInstance(id: string): Promise<void> {
@@ -260,8 +261,9 @@ export async function updateProviderInstance(
providerName: string,
apiKey?: string,
modelName?: string,
type: "generative" | "embedding" = "generative",
): Promise<void> {
ProviderManager.update(id, name, providerName, apiKey, modelName);
ProviderManager.update(id, name, providerName, apiKey, modelName, type);
}
export async function getProviderMappings(): Promise<Record<string, string>> {
@@ -275,6 +277,10 @@ export async function setProviderMapping(
ProviderManager.setMapping(task, providerInstanceId);
}
export async function getAvailableProviders(): Promise<LLMProviderMeta[]> {
export async function getAvailableProviders(): Promise<ModelProviderMeta[]> {
return AVAILABLE_PROVIDERS;
}
export async function regenerateEmbeddings(newProviderInstanceId?: string): Promise<void> {
await simulationManager.regenerateAllEmbeddings(newProviderInstanceId);
}

View File

@@ -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<SimSnapshot> {
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<void> {
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,

View File

@@ -30,46 +30,55 @@ export interface LLMCallRecord {
export interface ILLMProvider {
providerName: string;
// We use Zod to ensure the generic T matches the schema
generateStructuredResponse<T extends z.ZodTypeAny>(
request: LLMRequest<T>,
): Promise<LLMResponse<z.infer<T>>>;
lastCalls?: LLMCallRecord[];
}
export interface LLMProviderInstance {
export interface IEmbeddingProvider {
providerName: string;
embed(text: string): Promise<number[]>;
}
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",
},
];

View File

@@ -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;

View File

@@ -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<number[]> {
return this.model.embedQuery(text);
}
}

View File

@@ -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<number[]> {
// 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;
}
}

View File

@@ -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) {

View File

@@ -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);
});
});