Skip to content
Merged
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
2 changes: 1 addition & 1 deletion package.json
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
{
"name": "@hasna/gateway",
"version": "0.1.5",
"version": "0.1.6",
"description": "Open-source AI model gateway core for one-key, multi-provider routing across Hasna apps and self-hosted deployments",
"type": "module",
"main": "dist/index.js",
Expand Down
4 changes: 2 additions & 2 deletions src/budget.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import { createHash } from "node:crypto";
import type { GatewayBudgetConfig, GatewayConfig, GatewayUsage, OpenAIChatCompletionRequest } from "./types";
import type { GatewayBudgetConfig, GatewayConfig, GatewayRoutableRequest, GatewayUsage } from "./types";
import { GatewayHttpError } from "./errors";
import { hasUsageLedgerBackend, readBudgetLedgerRecords } from "./storage";

Expand Down Expand Up @@ -64,7 +64,7 @@ export function fingerprintGatewayKey(value: string | undefined): string | undef
}

export function budgetContextFromRequest(
request: OpenAIChatCompletionRequest,
request: GatewayRoutableRequest,
base: GatewayBudgetContext = {},
): GatewayBudgetContext {
return {
Expand Down
15 changes: 10 additions & 5 deletions src/config.ts
Original file line number Diff line number Diff line change
Expand Up @@ -792,12 +792,17 @@ function validateProductionRouteReadiness(config: GatewayConfig, env: Record<str
const unavailableRouteIds = config.routes
.filter((route) => {
const model = route.modelAliases?.[0] ?? route.id;
try {
resolveRoute({ config, env }, routeProbeRequest(model));
return false;
} catch {
return true;
// A route is ready if a keyed provider can satisfy it for either chat or
// embeddings, so embeddings-only routes are not probed as chat requests.
for (const operation of ["chat", "embeddings"] as const) {
try {
resolveRoute({ config, env }, routeProbeRequest(model), { operation });
return false;
} catch {
// Try the next operation before marking the route unavailable.
}
}
return true;
})
.map((route) => route.id);

Expand Down
165 changes: 163 additions & 2 deletions src/gateway.ts
Original file line number Diff line number Diff line change
Expand Up @@ -15,10 +15,12 @@ import { transformOpenAICompatibleStream } from "./streaming";
import type {
GatewayRouteCandidate,
GatewayRouteDecision,
GatewayRoutableRequest,
GatewayRuntimeOptions,
OpenAIChatCompletionRequest,
OpenAIEmbeddingsRequest,
} from "./types";
import { estimateCostUsd, normalizeUsage, toOpenAIUsage } from "./usage";
import { estimateCostUsd, normalizeUsage, toOpenAIEmbeddingsUsage, toOpenAIUsage } from "./usage";

type CompletionResult = {
body: Record<string, unknown>;
Expand Down Expand Up @@ -59,7 +61,7 @@ function metadataFor(
};
}

function includeGatewayMetadata(options: GatewayRuntimeOptions, request: OpenAIChatCompletionRequest): boolean {
function includeGatewayMetadata(options: GatewayRuntimeOptions, request: GatewayRoutableRequest): boolean {
if (request.gateway?.strict_openai_compatibility) return false;
return request.gateway?.include_gateway_metadata ?? options.config.server.includeGatewayMetadata;
}
Expand Down Expand Up @@ -267,6 +269,31 @@ async function callProvider(
});
}

async function callEmbeddingProvider(
options: GatewayRuntimeOptions,
request: OpenAIEmbeddingsRequest,
candidate: GatewayRouteCandidate,
): Promise<Response> {
const adapter = adapterForProvider(candidate.provider);
if (!adapter.embed) {
throw new GatewayHttpError({
status: 400,
type: "gateway_config_error",
code: "provider_embeddings_unsupported",
message: `Provider ${candidate.provider.id} does not support embeddings requests.`,
provider: candidate.provider.id,
});
}
return adapter.embed({
provider: candidate.provider,
model: candidate.model,
request,
apiKey: apiKeyFor(candidate, options.env ?? process.env),
timeoutMs: options.config.server.requestTimeoutMs,
fetchImpl: options.fetchImpl,
});
}

async function openProviderStream(
options: GatewayRuntimeOptions,
request: OpenAIChatCompletionRequest,
Expand Down Expand Up @@ -448,6 +475,140 @@ export async function createChatCompletion(
});
}

export async function createEmbeddings(
options: GatewayRuntimeOptions,
request: OpenAIEmbeddingsRequest,
): Promise<CompletionResult> {
const env = options.env ?? process.env;
const requestBudgetContext = budgetContextFromRequest(request, options.budgetContext);
await assertBudgetPreflight(options.config, requestBudgetContext, { env });
const route = resolveRoute(options, request, { operation: "embeddings" });
const maxAttempts = Math.min(options.config.server.maxFallbackAttempts, route.candidates.length);
let lastError: GatewayHttpError | undefined;

for (const candidate of route.candidates.slice(0, maxAttempts)) {
const started = Date.now();
const budgetContext = { ...requestBudgetContext, selectedModel: candidate.model.id };
try {
await assertBudgetPreflight(options.config, budgetContext, { env });
const response = await callEmbeddingProvider(options, request, candidate);
const latencyMs = Date.now() - started;

if (!response.ok) {
const error = await providerErrorFromResponse(candidate, response);
route.decision.attempts.push({
provider: candidate.provider.id,
model: candidate.model.id,
providerModel: candidate.model.providerModel,
status: "failed",
reason: error.message,
errorType: error.type,
errorCode: error.code,
retryable: error.retryable,
latencyMs,
});
lastError = error;
if (error.retryable) continue;
throw error;
}

route.decision.selected = candidate.model.id;
route.decision.attempts.push({
provider: candidate.provider.id,
model: candidate.model.id,
providerModel: candidate.model.providerModel,
status: "selected",
latencyMs,
});

const providerJson = await parseProviderJson(response);
const rawUsage = providerJson.usage;
if (rawUsage === undefined && options.rateLimit?.requiresStreamingUsage === true) {
throw new GatewayHttpError({
status: 429,
type: "gateway_rate_limit_error",
code: "gateway_token_usage_missing",
message: "Provider response did not include usage required to enforce a token rate limit.",
raw: { context: budgetContext },
});
}
const usage = normalizeUsage(rawUsage);
await options.rateLimit?.onUsage?.(usage);
const estimatedCostUsd = estimateCostUsd(usage, candidate.model);
const budgets = await evaluateBudgetPostflight(
options.config,
budgetContext,
spendFromUsage(usage, estimatedCostUsd),
{ env },
);
const body: Record<string, unknown> = {
...providerJson,
object: providerJson.object ?? "list",
model: candidate.model.id,
usage: toOpenAIEmbeddingsUsage(usage),
};

if (includeGatewayMetadata(options, request)) {
body.gateway = metadataFor(candidate, route.decision, estimatedCostUsd, budgets);
}

await appendUsageLedgerBestEffort({
config: options.config,
provider: candidate.provider,
model: candidate.model,
decision: route.decision,
context: budgetContext,
usage,
estimatedCostUsd,
budgets,
status: "success",
});

assertBudgetPostflight(budgets);
return { body, status: 200, decision: route.decision };
} catch (error) {
const gatewayError =
error instanceof GatewayHttpError
? error
: new GatewayHttpError({
status: 502,
type: "provider_unavailable",
code: "provider_fetch_failed",
message: error instanceof Error ? error.message : "Provider fetch failed.",
retryable: true,
provider: candidate.provider.id,
});

lastError = gatewayError;
if (!route.decision.attempts.some((attempt) => attempt.provider === candidate.provider.id && attempt.status === "failed")) {
route.decision.attempts.push({
provider: candidate.provider.id,
model: candidate.model.id,
providerModel: candidate.model.providerModel,
status: "failed",
reason: gatewayError.message,
errorType: gatewayError.type,
errorCode: gatewayError.code,
retryable: gatewayError.retryable,
latencyMs: Date.now() - started,
});
}

if (gatewayError.retryable) continue;
throw gatewayError;
}
}

throw new GatewayHttpError({
status: lastError?.status ?? 502,
type: lastError?.type ?? "gateway_routing_error",
code: lastError?.code ?? "all_routes_failed",
message: lastError?.message ?? `All route attempts failed for model '${request.model}'.`,
retryable: false,
raw: route.decision,
});
}

export async function createChatCompletionStream(
options: GatewayRuntimeOptions,
request: OpenAIChatCompletionRequest,
Expand Down
8 changes: 7 additions & 1 deletion src/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@ export {
validateRuntimeSecrets,
} from "./config";
export { GatewayHttpError, gatewayErrorResponse, jsonError } from "./errors";
export { createChatCompletion, createChatCompletionStream } from "./gateway";
export { createChatCompletion, createChatCompletionStream, createEmbeddings } from "./gateway";
export { appendUsageLedger } from "./ledger";
export { toCapabilityCard, toCapabilityCards, toCostEstimate, toDecisionEnvelope } from "./lib/contracts";
export {
Expand Down Expand Up @@ -53,6 +53,7 @@ export type {
GatewayRateLimitConfig,
GatewayRequestOptions,
GatewayResponseCacheConfig,
GatewayRoutableRequest,
GatewayRouteAttempt,
GatewayRouteCandidate,
GatewayRouteDecision,
Expand All @@ -65,8 +66,13 @@ export type {
GatewayServerConfigInput,
GatewayUsage,
OpenAIChatCompletionRequest,
OpenAIEmbeddingsInput,
OpenAIEmbeddingsRequest,
OpenAIEmbeddingsUsage,
OpenAIUsage,
ProviderAdapter,
ProviderBuildInput,
ProviderBuildBaseInput,
ProviderEmbeddingsBuildInput,
ProviderHttpRequest,
} from "./types";
2 changes: 1 addition & 1 deletion src/providers/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -30,4 +30,4 @@ export {
toOpenAIChatCompletionResponse,
} from "./anthropic";
export { GoogleGeminiAdapter, googleGeminiOpenAIBaseUrl } from "./google-gemini";
export { OpenAICompatibleAdapter, toProviderChatBody } from "./openai-compatible";
export { OpenAICompatibleAdapter, toProviderChatBody, toProviderEmbeddingsBody } from "./openai-compatible";
56 changes: 55 additions & 1 deletion src/providers/openai-compatible.ts
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,10 @@ import type {
GatewayProviderConfig,
GatewayProviderError,
OpenAIChatCompletionRequest,
OpenAIEmbeddingsRequest,
ProviderAdapter,
ProviderBuildInput,
ProviderEmbeddingsBuildInput,
ProviderHttpRequest,
} from "../types";

Expand Down Expand Up @@ -57,6 +59,14 @@ const openRouterProviderFields = new Set([

const vercelGatewayFields = new Set(["models", "order", "only", "caching", "providerTimeouts"]);

const embeddingsForwardedFields = new Set([
"model",
"input",
"encoding_format",
"dimensions",
"user",
]);

function joinUrl(baseUrl: string, path: string): string {
return `${baseUrl.replace(/\/+$/, "")}/${path.replace(/^\/+/, "")}`;
}
Expand Down Expand Up @@ -241,6 +251,18 @@ export function toProviderChatBody(
return body;
}

export function toProviderEmbeddingsBody(request: OpenAIEmbeddingsRequest, providerModel: string): Record<string, unknown> {
const body: Record<string, unknown> = {};
for (const [key, value] of Object.entries(request)) {
if (embeddingsForwardedFields.has(key) && value !== undefined) {
body[key] = value;
}
}

body.model = providerModel;
return body;
}

function createAbortSignal(timeoutMs: number, signal?: AbortSignal): AbortSignal {
if (signal) return signal;
return AbortSignal.timeout(timeoutMs);
Expand All @@ -249,7 +271,7 @@ function createAbortSignal(timeoutMs: number, signal?: AbortSignal): AbortSignal
export class OpenAICompatibleAdapter implements ProviderAdapter {
readonly id = "openai-compatible";
readonly kind = "openai-compatible";
readonly supports: GatewayModelCapability[] = ["chat", "streaming", "tools", "json"];
readonly supports: GatewayModelCapability[] = ["chat", "streaming", "tools", "json", "embeddings"];

buildRequest(input: ProviderBuildInput): ProviderHttpRequest {
const baseUrl = providerBaseUrl(input.provider, input.env);
Expand Down Expand Up @@ -277,6 +299,33 @@ export class OpenAICompatibleAdapter implements ProviderAdapter {
};
}

buildEmbeddingsRequest(input: ProviderEmbeddingsBuildInput): ProviderHttpRequest {
if (!input.provider.baseUrl) {
throw new Error(`Provider ${input.provider.id} does not define a baseUrl.`);
}

const body = toProviderEmbeddingsBody(input.request, input.model.providerModel);

return {
url: joinUrl(input.provider.baseUrl, "/embeddings"),
init: {
method: "POST",
headers: {
"content-type": "application/json",
authorization: `Bearer ${input.apiKey}`,
...(input.provider.id === "openrouter"
? {
"http-referer": "https://github.com/hasna/open-gateway",
"x-title": "Hasna Gateway",
}
: {}),
},
body: JSON.stringify(body),
signal: createAbortSignal(input.timeoutMs, input.signal),
},
};
}

send(input: ProviderBuildInput): Promise<Response> {
const request = this.buildRequest({
...input,
Expand All @@ -299,6 +348,11 @@ export class OpenAICompatibleAdapter implements ProviderAdapter {
return (input.fetchImpl ?? fetch)(request.url, request.init);
}

embed(input: ProviderEmbeddingsBuildInput): Promise<Response> {
const request = this.buildEmbeddingsRequest(input);
return (input.fetchImpl ?? fetch)(request.url, request.init);
}

mapError(response: Response, bodyText?: string): GatewayProviderError {
const mapped = mapProviderStatus(response.status);
return {
Expand Down
Loading
Loading