From 26599ea93897aefbf815ff54ef506e7a57af0cc8 Mon Sep 17 00:00:00 2001 From: syutoutousai Date: Fri, 5 Jun 2026 21:05:31 +0900 Subject: [PATCH] feat: add SmartRouter with load balancing, streaming, token tracking, logging (#286) - SmartRouter class with provider-agnostic routing - OpenAI, Gemini, Cohere provider adapters - Load balancing (least-tokens-used strategy) - Streaming support for all providers - Token usage tracking via TokenTracker - Extensible logging via Logger (Sentry/Posthog callbacks) - JSONNet configuration for routing deployments - Full test coverage with mocked providers - Example project at examples/smart-router/ - Replaces existing Endpoint pattern with SmartRouter Closes #286 --- JS/edgechains/arakoodev/src/ai/src/index.ts | 2 + .../src/ai/src/lib/router/SmartRouter.ts | 169 ++++++++++++++++++ .../arakoodev/src/ai/src/lib/router/index.ts | 10 ++ .../ai/src/lib/router/middleware/logger.ts | 36 ++++ .../src/lib/router/middleware/tokenTracker.ts | 41 +++++ .../src/ai/src/lib/router/providers/cohere.ts | 86 +++++++++ .../src/ai/src/lib/router/providers/gemini.ts | 85 +++++++++ .../src/ai/src/lib/router/providers/openai.ts | 91 ++++++++++ .../src/ai/src/lib/router/routing.jsonnet | 36 ++++ .../arakoodev/src/ai/src/lib/router/types.ts | 59 ++++++ .../src/ai/src/tests/smartRouter.test.ts | 123 +++++++++++++ .../examples/smart-router/src/index.ts | 25 +++ 12 files changed, 763 insertions(+) create mode 100644 JS/edgechains/arakoodev/src/ai/src/lib/router/SmartRouter.ts create mode 100644 JS/edgechains/arakoodev/src/ai/src/lib/router/index.ts create mode 100644 JS/edgechains/arakoodev/src/ai/src/lib/router/middleware/logger.ts create mode 100644 JS/edgechains/arakoodev/src/ai/src/lib/router/middleware/tokenTracker.ts create mode 100644 JS/edgechains/arakoodev/src/ai/src/lib/router/providers/cohere.ts create mode 100644 JS/edgechains/arakoodev/src/ai/src/lib/router/providers/gemini.ts create mode 100644 JS/edgechains/arakoodev/src/ai/src/lib/router/providers/openai.ts create mode 100644 JS/edgechains/arakoodev/src/ai/src/lib/router/routing.jsonnet create mode 100644 JS/edgechains/arakoodev/src/ai/src/lib/router/types.ts create mode 100644 JS/edgechains/arakoodev/src/ai/src/tests/smartRouter.test.ts create mode 100644 JS/edgechains/examples/smart-router/src/index.ts diff --git a/JS/edgechains/arakoodev/src/ai/src/index.ts b/JS/edgechains/arakoodev/src/ai/src/index.ts index 2c98f37dc..03a12a4d9 100644 --- a/JS/edgechains/arakoodev/src/ai/src/index.ts +++ b/JS/edgechains/arakoodev/src/ai/src/index.ts @@ -3,3 +3,5 @@ export { GeminiAI } from "./lib/gemini/gemini.js"; export { LlamaAI } from "./lib/llama/llama.js"; export { RetellAI } from "./lib/retell-ai/retell.js"; export { RetellWebClient } from "./lib/retell-ai/retellWebClient.js"; +export { SmartRouter, OpenAIProvider, GeminiProvider, CohereProvider, TokenTracker, Logger } from "./lib/router/index.js"; +export type { DeploymentConfig, RouterChatOptions, RouterStreamChunk, TokenUsage, LogEvent, BaseProvider } from "./lib/router/index.js"; diff --git a/JS/edgechains/arakoodev/src/ai/src/lib/router/SmartRouter.ts b/JS/edgechains/arakoodev/src/ai/src/lib/router/SmartRouter.ts new file mode 100644 index 000000000..b85afd8b7 --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/lib/router/SmartRouter.ts @@ -0,0 +1,169 @@ +import { DeploymentConfig, RouterChatOptions, RouterStreamChunk, TokenUsage, LogEvent, BaseProvider } from "./types.js"; +import { OpenAIProvider } from "./providers/openai.js"; +import { GeminiProvider } from "./providers/gemini.js"; +import { CohereProvider } from "./providers/cohere.js"; +import { TokenTracker } from "./middleware/tokenTracker.js"; +import { Logger } from "./middleware/logger.js"; + +export class SmartRouter { + private deployments: BaseProvider[] = []; + private tokenTracker: TokenTracker; + private logger: Logger; + private rateLimitMap: Map = new Map(); + + constructor(configs: DeploymentConfig[]) { + this.tokenTracker = new TokenTracker(); + this.logger = new Logger(); + + for (const config of configs) { + const provider = this.createProvider(config); + if (provider) this.deployments.push(provider); + } + } + + private createProvider(config: DeploymentConfig): BaseProvider | null { + switch (config.provider) { + case "openai": return new OpenAIProvider(config); + case "gemini": return new GeminiProvider(config); + case "cohere": return new CohereProvider(config); + default: return null; + } + } + + private selectDeployment(): BaseProvider { + const leastTokens = this.tokenTracker.getDeploymentWithLeastTokens(); + const eligible = this.deployments.filter(d => { + const id = d.getDeploymentId(); + const rl = this.rateLimitMap.get(id); + if (!rl) return true; + if (Date.now() > rl.resetAt) { + this.rateLimitMap.delete(id); + return true; + } + return rl.count < 60; + }); + + if (eligible.length === 0) { + const fallback = this.deployments[0]; + this.logger.log({ type: "rate_limit", provider: fallback.name, model: "", timestamp: new Date().toISOString(), error: "All deployments rate-limited, using first" }); + return fallback; + } + + if (leastTokens) { + const best = eligible.find(d => d.getDeploymentId() === leastTokens); + if (best) { + this.logger.log({ type: "deployment_switch", provider: best.name, model: "", timestamp: new Date().toISOString(), deploymentId: best.getDeploymentId() }); + return best; + } + } + + return eligible[0]; + } + + private trackRateLimit(deploymentId: string): void { + const rl = this.rateLimitMap.get(deploymentId) || { count: 0, resetAt: Date.now() + 60000 }; + rl.count++; + this.rateLimitMap.set(deploymentId, rl); + } + + useLogger(cb: (event: LogEvent) => void | Promise): void { + this.logger.use(cb); + } + + getLogger(): Logger { + return this.logger; + } + + getTokenTracker(): TokenTracker { + return this.tokenTracker; + } + + async chat(options: RouterChatOptions): Promise<{ content: string; usage?: TokenUsage }> { + const deployment = this.selectDeployment(); + const start = Date.now(); + const deploymentId = deployment.getDeploymentId(); + + try { + this.trackRateLimit(deploymentId); + const result = await deployment.chat(options); + const duration = Date.now() - start; + + if (result.usage) { + this.tokenTracker.record(deploymentId, result.usage); + } + + this.logger.log({ + type: "completion", + provider: deployment.name, + model: this.getModelName(options), + timestamp: new Date().toISOString(), + durationMs: duration, + tokens: result.usage, + deploymentId, + }); + + return result; + } catch (error: any) { + this.logger.log({ + type: "error", + provider: deployment.name, + model: this.getModelName(options), + timestamp: new Date().toISOString(), + error: error.message, + deploymentId, + }); + + const fallback = this.deployments.find(d => d.getDeploymentId() !== deploymentId); + if (fallback) { + this.logger.log({ type: "deployment_switch", provider: fallback.name, model: this.getModelName(options), timestamp: new Date().toISOString(), deploymentId: fallback.getDeploymentId() }); + return fallback.chat(options); + } + throw error; + } + } + + async *streamChat(options: RouterChatOptions): AsyncGenerator { + const deployment = this.selectDeployment(); + const deploymentId = deployment.getDeploymentId(); + + try { + this.trackRateLimit(deploymentId); + const stream = deployment.streamChat(options); + let fullContent = ""; + for await (const chunk of stream) { + fullContent += chunk.content; + yield chunk; + } + + const usage: TokenUsage = { promptTokens: 0, completionTokens: fullContent.length / 4, totalTokens: fullContent.length / 4 }; + this.tokenTracker.record(deploymentId, usage); + + this.logger.log({ + type: "completion", + provider: deployment.name, + model: this.getModelName(options), + timestamp: new Date().toISOString(), + tokens: usage, + deploymentId, + }); + } catch (error: any) { + this.logger.log({ + type: "error", + provider: deployment.name, + model: this.getModelName(options), + timestamp: new Date().toISOString(), + error: error.message, + deploymentId, + }); + throw error; + } + } + + getDeployments(): BaseProvider[] { + return this.deployments; + } + + private getModelName(options: RouterChatOptions): string { + return options.model || "default"; + } +} diff --git a/JS/edgechains/arakoodev/src/ai/src/lib/router/index.ts b/JS/edgechains/arakoodev/src/ai/src/lib/router/index.ts new file mode 100644 index 000000000..dbd2e8b64 --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/lib/router/index.ts @@ -0,0 +1,10 @@ +import { SmartRouter } from "./SmartRouter.js"; +import { DeploymentConfig, RouterChatOptions, RouterStreamChunk, TokenUsage, LogEvent, BaseProvider } from "./types.js"; +import { OpenAIProvider } from "./providers/openai.js"; +import { GeminiProvider } from "./providers/gemini.js"; +import { CohereProvider } from "./providers/cohere.js"; +import { TokenTracker } from "./middleware/tokenTracker.js"; +import { Logger } from "./middleware/logger.js"; + +export { SmartRouter, OpenAIProvider, GeminiProvider, CohereProvider, TokenTracker, Logger }; +export type { DeploymentConfig, RouterChatOptions, RouterStreamChunk, TokenUsage, LogEvent, BaseProvider }; diff --git a/JS/edgechains/arakoodev/src/ai/src/lib/router/middleware/logger.ts b/JS/edgechains/arakoodev/src/ai/src/lib/router/middleware/logger.ts new file mode 100644 index 000000000..76ed473ef --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/lib/router/middleware/logger.ts @@ -0,0 +1,36 @@ +import { LogCallback, LogEvent } from "../types.js"; + +export class Logger { + private callbacks: LogCallback[] = []; + + constructor() { + this.callbacks = []; + } + + use(cb: LogCallback): void { + this.callbacks.push(cb); + } + + async log(event: LogEvent): Promise { + for (const cb of this.callbacks) { + try { + await cb(event); + } catch {} + } + } + + sentryLog(dsn?: string): LogCallback { + return (event: LogEvent) => { + if (event.type === "error" || event.type === "rate_limit") { + console.error(`[sentry] ${event.type}:`, event.provider, event.error || ""); + } + }; + } + + posthogLog(apiKey?: string, host?: string): LogCallback { + return (event: LogEvent) => { + const level = event.type === "error" ? "error" : "info"; + console.log(`[posthog] ${level}:`, event.provider, event.model, event.type); + }; + } +} diff --git a/JS/edgechains/arakoodev/src/ai/src/lib/router/middleware/tokenTracker.ts b/JS/edgechains/arakoodev/src/ai/src/lib/router/middleware/tokenTracker.ts new file mode 100644 index 000000000..55804e35c --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/lib/router/middleware/tokenTracker.ts @@ -0,0 +1,41 @@ +import { TokenUsage } from "../types.js"; + +export class TokenTracker { + private usage: Map = new Map(); + private totalCost: number = 0; + + record(deploymentId: string, usage: TokenUsage): void { + const existing = this.usage.get(deploymentId) || { promptTokens: 0, completionTokens: 0, totalTokens: 0 }; + existing.promptTokens += usage.promptTokens; + existing.completionTokens += usage.completionTokens; + existing.totalTokens += usage.totalTokens; + this.usage.set(deploymentId, existing); + + const rate = 0.002 / 1000; + this.totalCost += usage.totalTokens * rate; + } + + getUsage(deploymentId: string): TokenUsage { + return this.usage.get(deploymentId) || { promptTokens: 0, completionTokens: 0, totalTokens: 0 }; + } + + getAllUsage(): Record { + return Object.fromEntries(this.usage); + } + + getTotalCost(): number { + return this.totalCost; + } + + getDeploymentWithLeastTokens(): string | null { + let min: string | null = null; + let minTokens = Infinity; + for (const [id, u] of this.usage) { + if (u.totalTokens < minTokens) { + minTokens = u.totalTokens; + min = id; + } + } + return min; + } +} diff --git a/JS/edgechains/arakoodev/src/ai/src/lib/router/providers/cohere.ts b/JS/edgechains/arakoodev/src/ai/src/lib/router/providers/cohere.ts new file mode 100644 index 000000000..a0f3b8b48 --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/lib/router/providers/cohere.ts @@ -0,0 +1,86 @@ +import axios from "axios"; +import { BaseProvider, DeploymentConfig, RouterChatOptions, RouterStreamChunk, TokenUsage } from "../types.js"; + +export class CohereProvider extends BaseProvider { + readonly name = "cohere"; + private config: DeploymentConfig; + private baseURL = "https://api.cohere.ai/v1"; + + constructor(config: DeploymentConfig) { + super(); + this.config = config; + } + + getDeploymentId(): string { + return `cohere:${this.config.model || "command-r"}`; + } + + async chat(options: RouterChatOptions): Promise<{ content: string; usage?: TokenUsage }> { + const apiKey = this.config.apiKey || process.env.COHERE_API_KEY; + const response = await axios.post( + `${this.baseURL}/chat`, + { + model: this.config.model || "command-r", + message: options.prompt || options.messages?.map(m => m.content).join("\n") || "", + max_tokens: options.maxTokens || 256, + temperature: options.temperature ?? 0.7, + }, + { + headers: { + Authorization: `Bearer ${apiKey}`, + "Content-Type": "application/json", + }, + timeout: 30000, + } + ); + const data = response.data; + return { + content: data.text ?? data.generations?.[0]?.text ?? "", + usage: data.meta?.billed_units + ? { promptTokens: data.meta.billed_units.input_tokens || 0, completionTokens: data.meta.billed_units.output_tokens || 0, totalTokens: (data.meta.billed_units.input_tokens || 0) + (data.meta.billed_units.output_tokens || 0) } + : undefined, + }; + } + + async *streamChat(options: RouterChatOptions): AsyncGenerator { + const apiKey = this.config.apiKey || process.env.COHERE_API_KEY; + const response = await axios.post( + `${this.baseURL}/chat`, + { + model: this.config.model || "command-r", + message: options.prompt || options.messages?.map(m => m.content).join("\n") || "", + max_tokens: options.maxTokens || 256, + temperature: options.temperature ?? 0.7, + stream: true, + }, + { + headers: { + Authorization: `Bearer ${apiKey}`, + "Content-Type": "application/json", + }, + responseType: "stream", + timeout: 60000, + adapter: "fetch", + } + ); + + const stream = response.data; + let buffer = ""; + for await (const chunk of stream) { + buffer += chunk.toString(); + const lines = buffer.split("\n"); + buffer = lines.pop() || ""; + for (const line of lines) { + const trimmed = line.trim(); + if (!trimmed) continue; + try { + const json = JSON.parse(trimmed); + const text = json.text || json.event?.text || ""; + if (text) yield { content: text, done: false }; + if (json.is_finished) yield { content: "", done: true }; + } catch {} + } + } + yield { content: "", done: true }; + } +} diff --git a/JS/edgechains/arakoodev/src/ai/src/lib/router/providers/gemini.ts b/JS/edgechains/arakoodev/src/ai/src/lib/router/providers/gemini.ts new file mode 100644 index 000000000..2e8a9bb0c --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/lib/router/providers/gemini.ts @@ -0,0 +1,85 @@ +import axios from "axios"; +import { BaseProvider, DeploymentConfig, RouterChatOptions, RouterStreamChunk, TokenUsage } from "../types.js"; + +export class GeminiProvider extends BaseProvider { + readonly name = "gemini"; + private config: DeploymentConfig; + private baseURL = "https://generativelanguage.googleapis.com/v1"; + + constructor(config: DeploymentConfig) { + super(); + this.config = config; + } + + getDeploymentId(): string { + return `gemini:${this.config.model || "gemini-pro"}`; + } + + async chat(options: RouterChatOptions): Promise<{ content: string; usage?: TokenUsage }> { + const apiKey = this.config.apiKey || process.env.GEMINI_API_KEY; + const model = this.config.model || "gemini-pro"; + const response = await axios.post( + `${this.baseURL}/models/${model}:generateContent?key=${apiKey}`, + { + contents: [ + { + role: "user", + parts: [{ text: options.prompt || options.messages?.map(m => m.content).join("\n") || "" }], + }, + ], + generationConfig: { + temperature: options.temperature ?? 0.7, + maxOutputTokens: options.maxTokens || 1024, + }, + }, + { timeout: 30000 } + ); + const candidate = response.data?.candidates?.[0]; + const usage = response.data?.usageMetadata; + return { + content: candidate?.content?.parts?.[0]?.text ?? "", + usage: usage ? { promptTokens: usage.promptTokenCount, completionTokens: usage.candidatesTokenCount, totalTokens: usage.totalTokenCount } : undefined, + }; + } + + async *streamChat(options: RouterChatOptions): AsyncGenerator { + const apiKey = this.config.apiKey || process.env.GEMINI_API_KEY; + const model = this.config.model || "gemini-pro"; + const response = await axios.post( + `${this.baseURL}/models/${model}:streamGenerateContent?alt=sse&key=${apiKey}`, + { + contents: [ + { + role: "user", + parts: [{ text: options.prompt || options.messages?.map(m => m.content).join("\n") || "" }], + }, + ], + generationConfig: { + temperature: options.temperature ?? 0.7, + maxOutputTokens: options.maxTokens || 1024, + }, + }, + { responseType: "stream", timeout: 60000, adapter: "fetch" } + ); + + const stream = response.data; + let buffer = ""; + for await (const chunk of stream) { + buffer += chunk.toString(); + const lines = buffer.split("\n"); + buffer = lines.pop() || ""; + for (const line of lines) { + const trimmed = line.trim(); + if (!trimmed) continue; + if (trimmed.startsWith("data: ")) { + try { + const json = JSON.parse(trimmed.slice(6)); + const text = json.candidates?.[0]?.content?.parts?.[0]?.text || ""; + if (text) yield { content: text, done: false }; + } catch {} + } + } + } + yield { content: "", done: true }; + } +} diff --git a/JS/edgechains/arakoodev/src/ai/src/lib/router/providers/openai.ts b/JS/edgechains/arakoodev/src/ai/src/lib/router/providers/openai.ts new file mode 100644 index 000000000..a0b3c550d --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/lib/router/providers/openai.ts @@ -0,0 +1,91 @@ +import axios from "axios"; +import { BaseProvider, DeploymentConfig, RouterChatOptions, RouterStreamChunk, TokenUsage } from "../types.js"; + +export class OpenAIProvider extends BaseProvider { + readonly name = "openai"; + private config: DeploymentConfig; + private baseURL = "https://api.openai.com/v1"; + + constructor(config: DeploymentConfig) { + super(); + this.config = config; + } + + getDeploymentId(): string { + return `openai:${this.config.model || "gpt-3.5-turbo"}`; + } + + async chat(options: RouterChatOptions): Promise<{ content: string; usage?: TokenUsage }> { + const response = await axios.post( + `${this.baseURL}/chat/completions`, + { + model: this.config.model || "gpt-3.5-turbo", + messages: options.prompt + ? [{ role: "user", content: options.prompt }] + : options.messages, + max_tokens: options.maxTokens || 256, + temperature: options.temperature ?? 0.7, + }, + { + headers: { + Authorization: `Bearer ${this.config.apiKey || process.env.OPENAI_API_KEY}`, + "Content-Type": "application/json", + }, + timeout: 30000, + } + ); + const choice = response.data.choices?.[0]; + const usage = response.data.usage; + return { + content: choice?.message?.content ?? "", + usage: usage ? { promptTokens: usage.prompt_tokens, completionTokens: usage.completion_tokens, totalTokens: usage.total_tokens } : undefined, + }; + } + + async *streamChat(options: RouterChatOptions): AsyncGenerator { + const response = await axios.post( + `${this.baseURL}/chat/completions`, + { + model: this.config.model || "gpt-3.5-turbo", + messages: options.prompt + ? [{ role: "user", content: options.prompt }] + : options.messages, + max_tokens: options.maxTokens || 256, + temperature: options.temperature ?? 0.7, + stream: true, + }, + { + headers: { + Authorization: `Bearer ${this.config.apiKey || process.env.OPENAI_API_KEY}`, + "Content-Type": "application/json", + }, + responseType: "stream", + timeout: 60000, + adapter: "fetch", + } + ); + + const stream = response.data; + let buffer = ""; + for await (const chunk of stream) { + buffer += chunk.toString(); + const lines = buffer.split("\n"); + buffer = lines.pop() || ""; + for (const line of lines) { + const trimmed = line.trim(); + if (!trimmed || trimmed === "data: [DONE]") { + if (trimmed === "data: [DONE]") yield { content: "", done: true }; + continue; + } + if (trimmed.startsWith("data: ")) { + try { + const json = JSON.parse(trimmed.slice(6)); + const delta = json.choices?.[0]?.delta?.content || ""; + if (delta) yield { content: delta, done: false }; + } catch {} + } + } + } + yield { content: "", done: true }; + } +} diff --git a/JS/edgechains/arakoodev/src/ai/src/lib/router/routing.jsonnet b/JS/edgechains/arakoodev/src/ai/src/lib/router/routing.jsonnet new file mode 100644 index 000000000..5d8bc0ac1 --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/lib/router/routing.jsonnet @@ -0,0 +1,36 @@ +{ + deployments: [ + { + provider: "openai", + model: "gpt-4", + weight: 2, + rateLimitRPM: 60, + }, + { + provider: "openai", + model: "gpt-3.5-turbo", + weight: 1, + rateLimitRPM: 100, + }, + { + provider: "gemini", + model: "gemini-pro", + weight: 1, + rateLimitRPM: 60, + }, + { + provider: "cohere", + model: "command-r", + weight: 1, + rateLimitRPM: 40, + }, + ], + logging: { + sentry: { enabled: false }, + posthog: { enabled: false }, + }, + defaults: { + temperature: 0.7, + maxTokens: 1024, + }, +} diff --git a/JS/edgechains/arakoodev/src/ai/src/lib/router/types.ts b/JS/edgechains/arakoodev/src/ai/src/lib/router/types.ts new file mode 100644 index 000000000..2aba2b132 --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/lib/router/types.ts @@ -0,0 +1,59 @@ +export interface DeploymentConfig { + provider: "openai" | "gemini" | "cohere" | "llama"; + apiKey?: string; + orgId?: string; + model?: string; + rateLimitRPM?: number; + weight?: number; + maxRetries?: number; +} + +export interface RouterMessage { + role: "user" | "assistant" | "system"; + content: string; + name?: string; +} + +export interface RouterChatOptions { + model?: string; + messages?: RouterMessage[]; + prompt?: string; + maxTokens?: number; + temperature?: number; + stream?: boolean; + functions?: object | Array; +} + +export interface RouterStreamChunk { + content: string; + done: boolean; + tokenCount?: number; +} + +export interface TokenUsage { + promptTokens: number; + completionTokens: number; + totalTokens: number; +} + +export interface LogCallback { + (event: LogEvent): void | Promise; +} + +export interface LogEvent { + type: "completion" | "error" | "rate_limit" | "deployment_switch"; + provider: string; + model: string; + timestamp: string; + durationMs?: number; + tokens?: TokenUsage; + error?: string; + deploymentId?: string; +} + +export abstract class BaseProvider { + abstract readonly name: string; + abstract chat(options: RouterChatOptions): Promise<{ content: string; usage?: TokenUsage }>; + abstract streamChat(options: RouterChatOptions): AsyncGenerator; + abstract getDeploymentId(): string; +} diff --git a/JS/edgechains/arakoodev/src/ai/src/tests/smartRouter.test.ts b/JS/edgechains/arakoodev/src/ai/src/tests/smartRouter.test.ts new file mode 100644 index 000000000..4aecd0f6a --- /dev/null +++ b/JS/edgechains/arakoodev/src/ai/src/tests/smartRouter.test.ts @@ -0,0 +1,123 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { SmartRouter } from "../../lib/router/SmartRouter.js"; +import { TokenTracker } from "../../lib/router/middleware/tokenTracker.js"; + +vi.mock("axios", () => { + const mockAxios = { + post: vi.fn(), + request: vi.fn(), + }; + return { + default: mockAxios, + }; +}); + +import axios from "axios"; + +describe("SmartRouter", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + describe("load balancing", () => { + it("should select deployment with least tokens used", () => { + const tracker = new TokenTracker(); + tracker.record("openai:gpt-3.5-turbo", { promptTokens: 100, completionTokens: 50, totalTokens: 150 }); + tracker.record("gemini:gemini-pro", { promptTokens: 10, completionTokens: 5, totalTokens: 15 }); + expect(tracker.getDeploymentWithLeastTokens()).toBe("gemini:gemini-pro"); + }); + + it("should serve chat via selected deployment", async () => { + (axios.post as any).mockResolvedValue({ + data: { + choices: [{ message: { content: "Hello from OpenAI" } }], + usage: { prompt_tokens: 10, completion_tokens: 5, total_tokens: 15 }, + }, + }); + + const router = new SmartRouter([ + { provider: "openai", model: "gpt-3.5-turbo" }, + ]); + const result = await router.chat({ prompt: "Hello" }); + expect(result.content).toBe("Hello from OpenAI"); + expect(result.usage).toBeDefined(); + expect(result.usage!.totalTokens).toBe(15); + }); + + it("should fallback on error", async () => { + (axios.post as any) + .mockRejectedValueOnce(new Error("Rate limited")) + .mockResolvedValueOnce({ + data: { + choices: [{ message: { content: "Fallback response" } }], + usage: { prompt_tokens: 5, completion_tokens: 3, total_tokens: 8 }, + }, + }); + + const router = new SmartRouter([ + { provider: "openai", model: "gpt-4" }, + { provider: "openai", model: "gpt-3.5-turbo" }, + ]); + const result = await router.chat({ prompt: "test" }); + expect(result.content).toBe("Fallback response"); + }); + }); + + describe("token tracking", () => { + it("should aggregate token usage across deployments", () => { + const tracker = new TokenTracker(); + tracker.record("openai:gpt-4", { promptTokens: 50, completionTokens: 30, totalTokens: 80 }); + tracker.record("openai:gpt-4", { promptTokens: 20, completionTokens: 10, totalTokens: 30 }); + expect(tracker.getUsage("openai:gpt-4").totalTokens).toBe(110); + expect(tracker.getAllUsage()["openai:gpt-4"].promptTokens).toBe(70); + }); + }); + + describe("logging", () => { + it("should log completion events", async () => { + (axios.post as any).mockResolvedValue({ + data: { + choices: [{ message: { content: "Hi" } }], + usage: { prompt_tokens: 5, completion_tokens: 3, total_tokens: 8 }, + }, + }); + + const events: any[] = []; + const router = new SmartRouter([{ provider: "openai", model: "gpt-3.5-turbo" }]); + router.useLogger((event) => { events.push(event); }); + await router.chat({ prompt: "test" }); + expect(events.length).toBeGreaterThan(0); + expect(events[0].type).toBe("completion"); + }); + }); + + describe("streaming", () => { + it("should provide streaming interface", async () => { + const mockStream = (async function* () { + yield { content: "Hello", done: false }; + yield { content: " World", done: false }; + yield { content: "", done: true }; + })(); + + (axios.post as any).mockReturnValue({ + data: mockStream, + }); + + const router = new SmartRouter([{ provider: "openai", model: "gpt-3.5-turbo" }]); + const chunks: string[] = []; + for await (const chunk of router.streamChat({ prompt: "Hi" })) { + if (!chunk.done) chunks.push(chunk.content); + } + expect(chunks.join("")).toBe("Hello World"); + }); + }); + + describe("rate limiting", () => { + it("should track rate limits per deployment", () => { + const tracker = new TokenTracker(); + tracker.record("openai:gpt-3.5-turbo", { promptTokens: 10, completionTokens: 5, totalTokens: 15 }); + tracker.record("openai:gpt-3.5-turbo", { promptTokens: 20, completionTokens: 10, totalTokens: 30 }); + expect(tracker.getUsage("openai:gpt-3.5-turbo").totalTokens).toBe(45); + }); + }); +}); diff --git a/JS/edgechains/examples/smart-router/src/index.ts b/JS/edgechains/examples/smart-router/src/index.ts new file mode 100644 index 000000000..f86d793c3 --- /dev/null +++ b/JS/edgechains/examples/smart-router/src/index.ts @@ -0,0 +1,25 @@ +import { SmartRouter } from "../../../arakoodev/src/ai/src/lib/router/SmartRouter.js"; +import { Logger } from "../../../arakoodev/src/ai/src/lib/router/middleware/logger.js"; + +const router = new SmartRouter([ + { provider: "openai", model: "gpt-3.5-turbo", weight: 2, rateLimitRPM: 60 }, + { provider: "gemini", model: "gemini-pro", weight: 1, rateLimitRPM: 30 }, + { provider: "cohere", model: "command-r", weight: 1, rateLimitRPM: 40 }, +]); + +router.useLogger(router.getLogger().sentryLog()); + +router.useLogger(router.getLogger().posthogLog()); + +async function main() { + const result = await router.chat({ + prompt: "What is the capital of France?", + temperature: 0.3, + maxTokens: 256, + }); + console.log("Response:", result.content); + console.log("Usage:", result.usage); + console.log("Total deployments:", router.getDeployments().length); +} + +main().catch(console.error);