Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
66 changes: 66 additions & 0 deletions backend/src/core/apiKeyProviders.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
export type ApiKeyProvider = string;
export type ApiKeySource = "user" | "env" | null;

type ProviderRecord = {
readonly envVars: readonly string[];
};

// Table-driven: env-var names live here, not in a switch statement.
// Adding a new provider is one registerApiKeyProvider() call — no edits here.
const _providerRegistry = new Map<string, ProviderRecord>([
["claude", { envVars: ["ANTHROPIC_API_KEY", "CLAUDE_API_KEY"] }],
["gemini", { envVars: ["GEMINI_API_KEY"] }],
["openai", { envVars: ["OPENAI_API_KEY"] }],
["openrouter", { envVars: ["OPENROUTER_API_KEY"] }],
["courtlistener", { envVars: ["COURTLISTENER_API_TOKEN"] }],
]);

/**
* Register a new API-key provider so that getUserApiKeyStatus() and
* getUserApiKeys() include it automatically.
*
* Call once from your provider setup file alongside registerProvider():
*
* registerApiKeyProvider("bedrock", ["AWS_ACCESS_KEY_ID"]);
* registerApiKeyProvider("ollama", []); // no key required
*/
export function registerApiKeyProvider(
provider: string,
envVars: readonly string[],
): void {
_providerRegistry.set(provider, { envVars });
}

/** Returns provider IDs in registration order. */
export function getRegisteredProviders(): readonly string[] {
return [..._providerRegistry.keys()];
}

export function isApiKeyProvider(value: string): boolean {
return _providerRegistry.has(value);
}

export function normalizeApiKeyProvider(value: string): string | null {
return _providerRegistry.has(value) ? value : null;
}

/**
* Returns the platform API key for provider from environment variables,
* or null when none of the provider's env vars are set.
*
* Table-driven: the env var names are declared in the provider registry above,
* not hard-coded per-provider in this function body.
*/
export function envApiKey(provider: string): string | null {
const record = _providerRegistry.get(provider);
if (!record) return null;
for (const varName of record.envVars) {
const val = process.env[varName]?.trim();
if (val) return val;
}
return null;
}

export function hasEnvApiKey(provider: string): boolean {
return !!envApiKey(provider);
}
97 changes: 97 additions & 0 deletions backend/src/lib/llm/__tests__/registry.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
import { describe, it, expect, beforeEach } from "vitest";
import {
registerProvider,
getRegisteredProvider,
findProviderForModel,
registeredProviderIds,
allRegisteredModels,
_resetRegistryForTesting,
type LLMProviderAdapter,
} from "../registry";

function makeAdapter(id: string, prefixes: string[], models: string[] = []): LLMProviderAdapter {
return {
id,
matchesModel: (m) => prefixes.some((p) => m.startsWith(p)),
stream: async () => ({ fullText: "" }),
complete: async () => "",
models: { main: models, mid: [], low: [] },
};
}

beforeEach(() => {
_resetRegistryForTesting();
});

describe("registerProvider / getRegisteredProvider", () => {
it("stores and retrieves an adapter by id", () => {
const adapter = makeAdapter("test", ["test-"]);
registerProvider(adapter);
expect(getRegisteredProvider("test")).toBe(adapter);
});

it("returns undefined for an unknown id", () => {
expect(getRegisteredProvider("unknown")).toBeUndefined();
});

it("re-registration replaces the previous adapter", () => {
const first = makeAdapter("p", ["p-"]);
const second = makeAdapter("p", ["p-"]);
registerProvider(first);
registerProvider(second);
expect(getRegisteredProvider("p")).toBe(second);
});
});

describe("findProviderForModel", () => {
it("returns the first provider whose matchesModel is true", () => {
const a = makeAdapter("alpha", ["alpha-"]);
const b = makeAdapter("beta", ["beta-"]);
registerProvider(a);
registerProvider(b);
expect(findProviderForModel("alpha-turbo")).toBe(a);
expect(findProviderForModel("beta-fast")).toBe(b);
});

it("returns undefined when no provider matches", () => {
registerProvider(makeAdapter("x", ["x-"]));
expect(findProviderForModel("unknown-model")).toBeUndefined();
});

it("the first registered provider wins on overlap", () => {
const first = makeAdapter("first", ["shared-"]);
const second = makeAdapter("second", ["shared-"]);
registerProvider(first);
registerProvider(second);
expect(findProviderForModel("shared-model")).toBe(first);
});
});

describe("registeredProviderIds", () => {
it("returns ids in insertion order", () => {
registerProvider(makeAdapter("c", ["c-"]));
registerProvider(makeAdapter("a", ["a-"]));
registerProvider(makeAdapter("b", ["b-"]));
expect(registeredProviderIds()).toEqual(["c", "a", "b"]);
});

it("returns an empty array when no providers are registered", () => {
expect(registeredProviderIds()).toEqual([]);
});
});

describe("allRegisteredModels", () => {
it("returns the union of all provider model lists", () => {
registerProvider(makeAdapter("p1", ["m-"], ["m1", "m2"]));
registerProvider(makeAdapter("p2", ["n-"], ["m2", "n1"]));
const set = allRegisteredModels();
expect(set.has("m1")).toBe(true);
expect(set.has("m2")).toBe(true);
expect(set.has("n1")).toBe(true);
expect(set.size).toBe(3);
});

it("returns an empty set when no providers are registered", () => {
expect(allRegisteredModels().size).toBe(0);
});
});
103 changes: 86 additions & 17 deletions backend/src/lib/llm/index.ts
Original file line number Diff line number Diff line change
@@ -1,30 +1,99 @@
import { streamClaude, completeClaudeText } from "./claude";
import { streamGemini, completeGeminiText } from "./gemini";
import { streamOpenAI, completeOpenAIText } from "./openai";
import { providerForModel } from "./models";
import type { StreamChatParams, StreamChatResult, UserApiKeys } from "./types";
import { registerProvider, getRegisteredProvider } from "./registry";
import {
providerForModel,
CLAUDE_MAIN_MODELS,
CLAUDE_MID_MODELS,
CLAUDE_LOW_MODELS,
GEMINI_MAIN_MODELS,
GEMINI_MID_MODELS,
GEMINI_LOW_MODELS,
OPENAI_MAIN_MODELS,
OPENAI_MID_MODELS,
OPENAI_LOW_MODELS,
} from "./models";
import type { StreamChatParams, StreamChatResult, CompleteTextParams } from "./types";

export * from "./types";
export * from "./models";

/**
* Register a third-party LLM provider so it is available via
* streamChatWithTools() and completeText().
*
* OpenAI-compatible providers can be added the same way — call
* registerProvider()/registerApiKeyProvider(), no core edits.
*/
export { registerProvider } from "./registry";
import { setupDemo } from "./providers/demo";

// ---------------------------------------------------------------------------
// Register built-in providers
// ---------------------------------------------------------------------------
// Providers are imported above so that Vitest's vi.mock() hoisting works:
// test files mock e.g. "../claude" before this module loads, so the mocked
// function is captured here and ends up in the registry.

/** Register the built-in LLM providers (claude/gemini/openai). */
export function registerBuiltinProviders(): void {
// Keyless demo model — always available. Lets a brand-new user get a
// response before any API key is configured, and backs the auto-fallback
// in routes/chat.ts.
setupDemo();

registerProvider({
id: "claude",
matchesModel: (m) => m.startsWith("claude"),
stream: streamClaude,
complete: completeClaudeText,
models: { main: CLAUDE_MAIN_MODELS, mid: CLAUDE_MID_MODELS, low: CLAUDE_LOW_MODELS },
});
registerProvider({
id: "gemini",
matchesModel: (m) => m.startsWith("gemini"),
stream: streamGemini,
complete: completeGeminiText,
models: { main: GEMINI_MAIN_MODELS, mid: GEMINI_MID_MODELS, low: GEMINI_LOW_MODELS },
});
registerProvider({
id: "openai",
matchesModel: (m) => m.startsWith("gpt-"),
stream: streamOpenAI,
complete: completeOpenAIText,
models: { main: OPENAI_MAIN_MODELS, mid: OPENAI_MID_MODELS, low: OPENAI_LOW_MODELS },
});
}

registerBuiltinProviders();

// ---------------------------------------------------------------------------
// Public dispatch
// ---------------------------------------------------------------------------

function requireAdapter(providerId: string, model: string) {
const adapter = getRegisteredProvider(providerId);
if (!adapter) {
throw new Error(
`LLM provider "${providerId}" matched model "${model}" but is not registered. ` +
`Import "lib/llm" to initialize built-in providers, ` +
`or call registerProvider() for third-party providers.`,
);
}
return adapter;
}

export async function streamChatWithTools(
params: StreamChatParams,
): Promise<StreamChatResult> {
const provider = providerForModel(params.model);
if (provider === "claude") return streamClaude(params);
if (provider === "openai") return streamOpenAI(params);
return streamGemini(params);
const providerId = providerForModel(params.model);
const adapter = requireAdapter(providerId, params.model);
return adapter.stream(params);
}

export async function completeText(params: {
model: string;
systemPrompt?: string;
user: string;
maxTokens?: number;
apiKeys?: UserApiKeys;
}): Promise<string> {
const provider = providerForModel(params.model);
if (provider === "claude") return completeClaudeText(params);
if (provider === "openai") return completeOpenAIText(params);
return completeGeminiText(params);
export async function completeText(params: CompleteTextParams): Promise<string> {
const providerId = providerForModel(params.model);
const adapter = requireAdapter(providerId, params.model);
return adapter.complete(params);
}
49 changes: 45 additions & 4 deletions backend/src/lib/llm/models.ts
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import type { Provider } from "./types";
import { findProviderForModel, allRegisteredModels } from "./registry";

// ---------------------------------------------------------------------------
// Canonical model IDs
// Canonical model IDs (built-in providers)
// ---------------------------------------------------------------------------
// Main-chat tier (top-end) — user picks one of these per message.
export const CLAUDE_MAIN_MODELS = [
Expand Down Expand Up @@ -32,6 +32,28 @@ export const DEFAULT_MAIN_MODEL = "gemini-3-flash-preview";
export const DEFAULT_TITLE_MODEL = "gemini-3.1-flash-lite-preview";
export const DEFAULT_TABULAR_MODEL = "gemini-3-flash-preview";

/**
* Built-in keyless "demo" model. Requires no API key and returns a canned,
* context-aware placeholder answer. Used as the automatic fallback when a
* request's chosen provider has no configured key, so a brand-new user still
* gets a response (and a nudge to add a real key) instead of a hard error.
* Also selectable directly in the model picker. Registered by
* providers/demo.ts.
*/
export const DEMO_MODEL = "mike-demo";

// Derived (not hand-maintained) fallback set for resolveModel().
// Built by spreading the *_MODELS arrays above, so adding a model to any
// of those arrays automatically includes it here — no second edit site.
//
// Why keep this alongside allRegisteredModels()? Two reasons:
// 1. Test isolation: models.test.ts imports models.ts directly without
// importing index.ts, so no providers are registered and the registry
// is empty. ALL_MODELS provides the fallback in that case.
// 2. External providers registered via registerProvider() appear in
// allRegisteredModels() but NOT here — that's intentional.
// resolveModel() checks both, so external models are always accepted
// once their provider is registered.
const ALL_MODELS = new Set<string>([
...CLAUDE_MAIN_MODELS,
...GEMINI_MAIN_MODELS,
Expand All @@ -48,14 +70,33 @@ const ALL_MODELS = new Set<string>([
// Provider inference
// ---------------------------------------------------------------------------

export function providerForModel(model: string): Provider {
/**
* Maps a model ID to its provider string.
*
* Registered providers are checked first so that externally registered
* adapters (Ollama, Bedrock, Azure) override the built-in prefix matching
* below — no edits to this file required to support a new provider.
*
* The prefix fallback keeps this function usable in test contexts that don't
* import index.ts and therefore don't trigger provider registration.
*/
export function providerForModel(model: string): string {
const registered = findProviderForModel(model);
if (registered) return registered.id;
if (model.startsWith("claude")) return "claude";
if (model.startsWith("gemini")) return "gemini";
if (model.startsWith("gpt-")) return "openai";
throw new Error(`Unknown model id: ${model}`);
}

/**
* Returns id if it is a recognised model, otherwise returns fallback.
*
* Checks the live registry first (includes externally registered models) then
* falls back to the static ALL_MODELS set so the function works in test
* contexts where no providers have been registered.
*/
export function resolveModel(id: string | null | undefined, fallback: string): string {
if (id && ALL_MODELS.has(id)) return id;
if (id && (allRegisteredModels().has(id) || ALL_MODELS.has(id))) return id;
return fallback;
}
Loading
Loading