mirror of
https://github.com/sortedcord/omnia.git
synced 2026-07-22 03:52:48 +05:30
major(llm): Support Embedding Providers
This commit is contained in:
@@ -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",
|
||||
},
|
||||
];
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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);
|
||||
});
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user