mirror of
https://github.com/simstudioai/sim.git
synced 2026-09-24 15:45:35 +08:00
feat(litellm): add LiteLLM as AI gateway provider (#4739)
* feat: add LiteLLM as AI gateway provider * fix: add litellm to attachments, provider store, utils, and block guards * fix: add frontend model discovery pipeline for litellm provider Add API route, contract, query hook case, and ProviderModelsLoader entry so litellm models are fetched and synced to the store on workspace load, matching the vllm/ollama/openrouter/fireworks pattern. Also fixes defaultModel to empty string and adds litellm/ prefix early-return in blocks/utils.ts (reviewer feedback). * fix: remove azureEndpoint fallback from LiteLLM provider Copy-paste artifact from vLLM provider. LiteLLM should only use LITELLM_BASE_URL, not fall back to azureEndpoint which could cause requests to be routed to the wrong server. * fix(litellm): close audit gaps from PR #4644 - byok.ts: add litellm branch to getApiKeyWithBYOK so workflow block execution can resolve the proxy key instead of throwing "API key is required for litellm ..." - check-api-validation-contracts.ts: bump route baseline 755 -> 756 to account for the new /api/providers/litellm/models route - .env.example: document LITELLM_BASE_URL / LITELLM_API_KEY - copilot edit-workflow validation: include LiteLLM in the list of user-configured prefixed providers shown to the model - providers/utils.ts: drop stray optional-chain on providers.litellm to match the vllm pattern - lint: apply biome formatting fixes (multi-line if, SVG path, multi-line DYNAMIC_MODEL_PROVIDERS) * fix(litellm): final parity gaps from second audit - blocks/utils.ts getModelOptions(): include litellm models in the combined model dropdown — was previously dropping any proxy-discovered models from the agent block model picker. - get-blocks-metadata-tool.ts mockProvidersState: add litellm bucket so the server-side copilot block-metadata fallback can render model options when the providers store is not initialized. - blocks/utils.test.ts: add litellm to mock providers state (initial + beforeEach reset) and add a parallel store-bucket guard test mirroring the vLLM case. - providers/utils.test.ts: add parallel getApiKey test for litellm. * feat(litellm): use official LiteLLM brand icon and color - icons.tsx: replace the placeholder letterform with the official LiteLLM brand mark embedded as a PNG data URI in an SVG image. - models.ts: set color: #040229 on the litellm provider definition to match the brand background. * chore(litellm): validate /v1/models response with shared schema in initialize() Match the API route handler — both code paths now run the same vllmUpstreamResponseSchema.parse() over the upstream /v1/models JSON instead of a raw type-cast, so malformed upstream payloads surface a descriptive ZodError instead of a downstream TypeError. Addresses Greptile review feedback on PR #4739. --------- Co-authored-by: RheagalFire <arishalam121@gmail.com>
This commit is contained in:
@@ -48,6 +48,8 @@ API_ENCRYPTION_KEY=your_api_encryption_key # Use `openssl rand -hex 32` to gener
|
||||
# OLLAMA_URL=http://localhost:11434 # URL for local Ollama server - uncomment if using local models
|
||||
# VLLM_BASE_URL=http://localhost:8000 # Base URL for your self-hosted vLLM (OpenAI-compatible)
|
||||
# VLLM_API_KEY= # Optional bearer token if your vLLM instance requires auth
|
||||
# LITELLM_BASE_URL=http://localhost:4000 # Base URL for your LiteLLM proxy (OpenAI-compatible)
|
||||
# LITELLM_API_KEY= # Optional bearer token if your LiteLLM proxy requires auth
|
||||
# FIREWORKS_API_KEY= # Optional Fireworks AI API key for model listing
|
||||
# NEXT_PUBLIC_BEDROCK_DEFAULT_CREDENTIALS=true # Set when using AWS default credential chain (IAM roles, ECS task roles, IRSA). Hides credential fields in Agent block UI.
|
||||
# AZURE_OPENAI_ENDPOINT= # Azure OpenAI endpoint (hides field in UI when set alongside NEXT_PUBLIC_AZURE_CONFIGURED)
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
import { createLogger } from '@sim/logger'
|
||||
import { getErrorMessage } from '@sim/utils/errors'
|
||||
import { type NextRequest, NextResponse } from 'next/server'
|
||||
import {
|
||||
providerModelsResponseSchema,
|
||||
vllmUpstreamResponseSchema,
|
||||
} from '@/lib/api/contracts/providers'
|
||||
import { env } from '@/lib/core/config/env'
|
||||
import { withRouteHandler } from '@/lib/core/utils/with-route-handler'
|
||||
import { filterBlacklistedModels, isProviderBlacklisted } from '@/providers/utils'
|
||||
|
||||
const logger = createLogger('LiteLLMModelsAPI')
|
||||
|
||||
export const GET = withRouteHandler(async (_request: NextRequest) => {
|
||||
if (isProviderBlacklisted('litellm')) {
|
||||
logger.info('LiteLLM provider is blacklisted, returning empty models')
|
||||
return NextResponse.json({ models: [] })
|
||||
}
|
||||
|
||||
const baseUrl = (env.LITELLM_BASE_URL || '').replace(/\/$/, '')
|
||||
|
||||
if (!baseUrl) {
|
||||
logger.info('LITELLM_BASE_URL not configured')
|
||||
return NextResponse.json({ models: [] })
|
||||
}
|
||||
|
||||
try {
|
||||
logger.info('Fetching LiteLLM models', { baseUrl })
|
||||
|
||||
const headers: Record<string, string> = {
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
|
||||
if (env.LITELLM_API_KEY) {
|
||||
headers.Authorization = `Bearer ${env.LITELLM_API_KEY}`
|
||||
}
|
||||
|
||||
const response = await fetch(`${baseUrl}/v1/models`, {
|
||||
headers,
|
||||
next: { revalidate: 60 },
|
||||
})
|
||||
|
||||
if (!response.ok) {
|
||||
logger.warn('LiteLLM service is not available', {
|
||||
status: response.status,
|
||||
statusText: response.statusText,
|
||||
})
|
||||
return NextResponse.json({ models: [] })
|
||||
}
|
||||
|
||||
const data = vllmUpstreamResponseSchema.parse(await response.json())
|
||||
const allModels = data.data.map((model) => `litellm/${model.id}`)
|
||||
const models = filterBlacklistedModels(allModels)
|
||||
|
||||
logger.info('Successfully fetched LiteLLM models', {
|
||||
count: models.length,
|
||||
filtered: allModels.length - models.length,
|
||||
models,
|
||||
})
|
||||
|
||||
return NextResponse.json(providerModelsResponseSchema.parse({ models }))
|
||||
} catch (error) {
|
||||
logger.error('Failed to fetch LiteLLM models', {
|
||||
error: getErrorMessage(error, 'Unknown error'),
|
||||
baseUrl,
|
||||
})
|
||||
|
||||
return NextResponse.json({ models: [] })
|
||||
}
|
||||
})
|
||||
@@ -6,6 +6,7 @@ import { useParams } from 'next/navigation'
|
||||
import { useProviderModels } from '@/hooks/queries/providers'
|
||||
import {
|
||||
updateFireworksProviderModels,
|
||||
updateLiteLLMProviderModels,
|
||||
updateOllamaProviderModels,
|
||||
updateOpenRouterProviderModels,
|
||||
updateVLLMProviderModels,
|
||||
@@ -32,6 +33,8 @@ function useSyncProvider(provider: ProviderName, workspaceId?: string) {
|
||||
updateOllamaProviderModels(data.models)
|
||||
} else if (provider === 'vllm') {
|
||||
updateVLLMProviderModels(data.models)
|
||||
} else if (provider === 'litellm') {
|
||||
updateLiteLLMProviderModels(data.models)
|
||||
} else if (provider === 'openrouter') {
|
||||
void updateOpenRouterProviderModels(data.models)
|
||||
if (data.modelInfo) {
|
||||
@@ -61,6 +64,7 @@ export function ProviderModelsLoader() {
|
||||
useSyncProvider('base')
|
||||
useSyncProvider('ollama')
|
||||
useSyncProvider('vllm')
|
||||
useSyncProvider('litellm')
|
||||
useSyncProvider('openrouter')
|
||||
useSyncProvider('fireworks', workspaceId)
|
||||
return null
|
||||
|
||||
@@ -27,6 +27,7 @@ const { mockProviders } = vi.hoisted(() => ({
|
||||
base: { models: [] as string[], isLoading: false },
|
||||
ollama: { models: [] as string[], isLoading: false },
|
||||
vllm: { models: [] as string[], isLoading: false },
|
||||
litellm: { models: [] as string[], isLoading: false },
|
||||
openrouter: { models: [] as string[], isLoading: false },
|
||||
fireworks: { models: [] as string[], isLoading: false },
|
||||
},
|
||||
@@ -101,6 +102,7 @@ describe('getApiKeyCondition / shouldRequireApiKeyForModel', () => {
|
||||
base: { models: [], isLoading: false },
|
||||
ollama: { models: [], isLoading: false },
|
||||
vllm: { models: [], isLoading: false },
|
||||
litellm: { models: [], isLoading: false },
|
||||
openrouter: { models: [], isLoading: false },
|
||||
fireworks: { models: [], isLoading: false },
|
||||
}
|
||||
@@ -185,6 +187,11 @@ describe('getApiKeyCondition / shouldRequireApiKeyForModel', () => {
|
||||
expect(evaluateCondition('my-custom-model')).toBe(false)
|
||||
})
|
||||
|
||||
it('does not require API key when model is in the LiteLLM store bucket', () => {
|
||||
mockProviders.value.litellm.models = ['litellm/anthropic/claude-sonnet-4-6']
|
||||
expect(evaluateCondition('litellm/anthropic/claude-sonnet-4-6')).toBe(false)
|
||||
})
|
||||
|
||||
it('requires API key when model is in the fireworks store bucket', () => {
|
||||
mockProviders.value.fireworks.models = ['fireworks/llama-3']
|
||||
expect(evaluateCondition('fireworks/llama-3')).toBe(true)
|
||||
|
||||
@@ -51,6 +51,7 @@ export function getModelOptions() {
|
||||
const baseModels = providersState.providers.base.models
|
||||
const ollamaModels = providersState.providers.ollama.models
|
||||
const vllmModels = providersState.providers.vllm.models
|
||||
const litellmModels = providersState.providers.litellm.models
|
||||
const openrouterModels = providersState.providers.openrouter.models
|
||||
const fireworksModels = providersState.providers.fireworks.models
|
||||
const allModels = Array.from(
|
||||
@@ -58,6 +59,7 @@ export function getModelOptions() {
|
||||
...baseModels,
|
||||
...ollamaModels,
|
||||
...vllmModels,
|
||||
...litellmModels,
|
||||
...openrouterModels,
|
||||
...fireworksModels,
|
||||
])
|
||||
@@ -160,12 +162,13 @@ function shouldRequireApiKeyForModel(model: string): boolean {
|
||||
) {
|
||||
return false
|
||||
}
|
||||
if (normalizedModel.startsWith('vllm/')) {
|
||||
if (normalizedModel.startsWith('vllm/') || normalizedModel.startsWith('litellm/')) {
|
||||
return false
|
||||
}
|
||||
|
||||
const storeProvider = getProviderFromStore(normalizedModel)
|
||||
if (storeProvider === 'ollama' || storeProvider === 'vllm') return false
|
||||
if (storeProvider === 'ollama' || storeProvider === 'vllm' || storeProvider === 'litellm')
|
||||
return false
|
||||
if (storeProvider) return true
|
||||
|
||||
if (isOllamaConfigured) {
|
||||
|
||||
@@ -4439,6 +4439,20 @@ export function VllmIcon(props: SVGProps<SVGSVGElement>) {
|
||||
)
|
||||
}
|
||||
|
||||
export function LitellmIcon(props: SVGProps<SVGSVGElement>) {
|
||||
return (
|
||||
<svg {...props} fill='none' viewBox='0 0 72 72' xmlns='http://www.w3.org/2000/svg'>
|
||||
<title>LiteLLM</title>
|
||||
<image
|
||||
href='data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAEgAAABICAMAAABiM0N1AAAC+lBMVEUAAABmZmXo6uvl5uQ+S3bd4uPl5t5DRk/n6ed6nbji4tYvPmuosLyIqcOApcAnJB86Rm8xLSpgc5Oguszi49zs7eakqaPq6uRNXY7T08M5SYc0QGsoJiRRT0mFhXZ0d3CnuMRHS1LZ2cr4+PIoNGM0MyslIR9KV341Mi0qO22WnrrD0tuLkqBpj7AqM1STrMdlZU+amXyfn4SyxtPn6ONOWXzj49urwtLs7OOeoZ/F09xMX5C6yNY0NkPDytJ5eWBteqQvQGxndpGarMBGRDs0MyqlpY5LSD+RkXZqc48wPWTZ2dGtraZETnHX18xXV0YhHxt9fnSvtclUaZ87OjFvb1f4+PJ5eV9zc1vk5Nj8/Pbq6t/m5trt7eRoaFL19u/i4tXd3c/6+vTz8+zv7+dsbFba2szT08Po6Nzf39J2dl3s7OHW1scmNWM4RnksPnR+fmNkZE7x8elml7rNzbw8SX2BgWWIiGwmNmcrPHAhLFNlZVDCwq1MX6ArO2soNlwnN2zQ0MBEW6VQVVhEQTk2MSyjo4YxQXdwnsDLy7nGxrJFXas2SIQdJ0fX18saIj1IRzsuKiWxz+amx+GQt9Z3pMZLZbtIYbO7u6RBV6A9UpWnjF+Ks9Fgk7itrZI2SYk+ToQwQXyWl3kzQG6FhWl7qMlpmr+7vaq+vqi2tqI7TYxHVIQvO2ZNVV6fwdtThatJXJqyspkzRYB+gHiPj3FUWV0dGxeDrMvNz8fIyLY/VZs6TpFEVY6bmnxUaHqSlHg8OTSQmrQ4RXI8RGGZvNdxjqZSZqRGWZXGycC9wb2xt7NmeK6BlKRjhKNOWX5+gXt7fXZPWHNsgL9Ra7yirrKqqp9RY56kpJOuk2A0PVXqq1RKS0ZTU0Nxmrhug55kfZVCUolJUmxfYFpCRUaTs85ajrR2krGUpa1LZK2RnaRaWkzppUXa29h6jMWFo7XAwbRaia5ecH6rkmMqMkg2Okayu9N2h5KUloeXiWRqamBFU3qbhFjglzLFgynBltUQAAAAVXRSTlMACgUOhiGTMS/98YYc/v7+h4X+29dmU0n59PPy2Rz9/fr59fPz7urMoZ2Zgzn++/jz8tC/tKmooH9xZl1VTPv079zZ08q1m5h+dV3l49/U0rGreG9fuSw7vQAABwlJREFUWMPs0z+M0lAcwHFzBHIQYRCwYSM6KrkBIoRcHHTUOJJU2tAW0j+hU+vowlaXDnIkDmU8DAuO4OJEk0JggzABywFx4WC52d/j8eeiIODiwjcBpvfh995rH5w6derU/+rMZrefE4TfH4s5oH8ybHbCH3sR8flcLhfL8rosC4HjB7ETjojPxfKykMt9WPfp2FkIx7uAAEBO5lmJISkxmc3SUPJI6Pnbbq/X63a75a+339l0Kg0pipLJZI/b1sPHr+bzH6vmZUqkcKTrmMO5jL52hmbfNg2b4tIB6FAlHvV6NLUIqWppmWpaFAkxjHQYZLuMet1IUUuaVoCucQWtZUmSBE8A/+QAxh4PX3hUFRAQqlXDMK6WGUZ1kuN59CDthc7On3ndfaQAAkS9XqlU8rjKlVH1dAQZEgL7GY9aVJECCBCfF+UrN27nONSGzKYA7YHWDFKWCCLM9vARNBwMzFarkxNgqL9B9g2DlXx9YQzv7lo/O82GVavpus6zDMq1+6biYfeGyddvPOP2YNCaNBs1XYJXQ0zisjhq1+EQ4Qt1VCxp13BFaDMmDNGo8QyFltK41W8GonccjiPoHI0WZ1zojxcGS4rr/8eLlXQ6lUhwqETi6dZd+X1BGEfVNK3vnDQtnbm3ExoAWI/i0FcKSgC07cojL6ezGbwJ/WmnsToP+ICBCQ73XuRZHiVTAG3Z1ZvGxAz1ndOgxTMkBQwgyMAIAKtYGk+WygL0566Ej+Uvt5YFd8OQJAUQIDAJNu7HMSmOQ046qSj0b+P8ar7OQ5oM4wCOv7NjfySoUSJ5ZXhGdEFBdN8HRKdtYpnSYW7tSFrNpducW1LOzH90M8i5BelaUOug3Ky1wKiBHdJsFeEBYakEVnRDv+f3vDu0Wf/2wcF42fN9n+d5977MhTNuGwylRwsPF0AGOrlkR7ASUthLHmZwoQqFxQDe5BzanTNid6ZOKW9uKT2FGViVECvBmcD54fT7Pw8NPSM8k++x3no8K0K/yRNnnDY0nyo8QzI4F7y8/sLefcX79whzhZ876+rq6uurqjSmsvb29jLC2TgtdFlHmw2nyyFTlBOoELvYBFw6MPisDktVGhryer1lXqc1EIpYsvpRS2n54YKCohwh7i6CCDYOEfhQnQMV2jGayvysen9o3PYVHc9PHT4IGbxIezECEyEJGM7Kz1/WCRVcmLFVdqWs7AoRDEXOnCy59BAzxbDBAGYizGUT+Sx4qh4cqqcZ6DSIH0DkAbhisk9jO3FNjna4nfbg3mAkkChCBdT0/ipAMq0NDreTZJzObqesUozrmhlnb61qv5WLN2LxnlwsBIcfBPTJfMujAUaYTYPD2mgxOaHSbTJ1iyXx2HE2GqvqPYPkaZC7Ox8DdCiMRYWse0bQSjONepsMIiaZTGaqVEIoYlEc6XR6BoWH8ougcCY4tBzhjwTikbcBOBwOq7VRb29SiWUyMXDLLPJUhomlnU7PrQIcHzLyNnGaKgVvHFZIwFz0kKmUKCwQcbstFrFSlMpMWKVv1fT392u8H3B0YGgpamYZwDsrJuz2pibIlCgVNguyuZWqVGaj3mHUaOAaxLWwYy8GNRtaiKvofAcESEIiKSlRq+WiapUNiSwixWwmSw/LJqe69BxO2mLoOhfi5NXzQV/udBBPqZ6enk+++1RPz/0UJstuorP93kXOevXrr4GXaGBg4OePrpMBbUmPwbWgJ73JcOAJSk5mspSSykoyWfcX+vmv54aHh1+gc20n25AZHJitUilCVFfX1FSjmtmKmhpmkkqhEhH3k8zgrrnNfNfPbL4ecDy79wargkpx9fX1vSIqUvpSGO5KrCoUvuTjI11gHQNSqfTYaNI8voAnEPB4Ani/nGEWx8McwcdvPGkYPCAA/D/l8XlIIOVn8xkorVRgqDebzxcgHkWHZwcd+AN7OC9vLgNmZVgU1RCCw4j2+HRg3t+cOAF/CEIgckeGW+X7JuCFEECHfjjEkTFBiKYWL8hISgru8TEpj6wJXv8QWGPob+hF69bNnz9/3rx55uvHpbhHuE2j4SaOIlg++j+eyMhZsbGxi7b1Vly+DF+ZywEVfu/Ba3AWaMlLe1a7ZcMGJixuhkokVyuBXC5CKj+fVqt1uXS62qioqJt+WznhO5wFNrlaXUJvcoQx7H0is1iTuX5NGsTQzaiba8cIbYzHjkQCIXVISC6ykY5rfSInghNTqwOkpouKCd+ZtZl2SAkog2w+l/aVNo2Ly0+vxSXqXLXp3DEWxnYqgZglQ+4+2GRXDF1ITJROS8CEIsKGlqZiiJZCM2KxD65QWmYie8KYdJiOrjYdwuFwstQlADOkYLoUEJ+QlrB2Z2AYZ2lmQkJC5tIwHcTdNImKjo4eHxANNnG53BGjOImJiRzmP/Ubf4ltRM+YtyIAAAAASUVORK5CYII='
|
||||
width='72'
|
||||
height='72'
|
||||
preserveAspectRatio='xMidYMid meet'
|
||||
/>
|
||||
</svg>
|
||||
)
|
||||
}
|
||||
|
||||
export function PosthogIcon(props: SVGProps<SVGSVGElement>) {
|
||||
return (
|
||||
<svg
|
||||
|
||||
@@ -5,6 +5,7 @@ import { requestJson } from '@/lib/api/client/request'
|
||||
import {
|
||||
getBaseProviderModelsContract,
|
||||
getFireworksProviderModelsContract,
|
||||
getLitellmProviderModelsContract,
|
||||
getOllamaProviderModelsContract,
|
||||
getOpenRouterProviderModelsContract,
|
||||
getVllmProviderModelsContract,
|
||||
@@ -54,6 +55,8 @@ async function requestProviderModels(
|
||||
return requestJson(getOllamaProviderModelsContract, { signal })
|
||||
case 'vllm':
|
||||
return requestJson(getVllmProviderModelsContract, { signal })
|
||||
case 'litellm':
|
||||
return requestJson(getLitellmProviderModelsContract, { signal })
|
||||
case 'openrouter':
|
||||
return requestJson(getOpenRouterProviderModelsContract, { signal })
|
||||
case 'fireworks':
|
||||
|
||||
@@ -74,6 +74,12 @@ export async function getApiKeyWithBYOK(
|
||||
return { apiKey: userProvidedKey || env.VLLM_API_KEY || 'empty', isBYOK: false }
|
||||
}
|
||||
|
||||
const isLitellmModel =
|
||||
provider === 'litellm' || useProvidersStore.getState().providers.litellm.models.includes(model)
|
||||
if (isLitellmModel) {
|
||||
return { apiKey: userProvidedKey || env.LITELLM_API_KEY || 'empty', isBYOK: false }
|
||||
}
|
||||
|
||||
const isFireworksModel =
|
||||
provider === 'fireworks' ||
|
||||
useProvidersStore.getState().providers.fireworks.models.includes(model)
|
||||
|
||||
@@ -207,6 +207,15 @@ export const getOpenRouterProviderModelsContract = defineRouteContract({
|
||||
},
|
||||
})
|
||||
|
||||
export const getLitellmProviderModelsContract = defineRouteContract({
|
||||
method: 'GET',
|
||||
path: '/api/providers/litellm/models',
|
||||
response: {
|
||||
mode: 'json',
|
||||
schema: providerModelsResponseSchema,
|
||||
},
|
||||
})
|
||||
|
||||
export const getFireworksProviderModelsContract = defineRouteContract({
|
||||
method: 'GET',
|
||||
path: '/api/providers/fireworks/models',
|
||||
|
||||
@@ -768,6 +768,7 @@ function callOptionsWithFallback(
|
||||
base: { models: staticModels.map((m) => m.id) },
|
||||
ollama: { models: [] },
|
||||
vllm: { models: [] },
|
||||
litellm: { models: [] },
|
||||
openrouter: { models: [] },
|
||||
fireworks: { models: [] },
|
||||
},
|
||||
|
||||
@@ -369,7 +369,7 @@ export function validateValueForSubBlockType(
|
||||
blockType,
|
||||
field: fieldName,
|
||||
value,
|
||||
error: `Unknown model id "${trimmed}" for block "${blockType}". Read components/blocks/${blockType}.json (the model.options array) for valid ids; prefer entries with recommended: true and avoid deprecated: true. For user-configured models (Ollama, vLLM, OpenRouter, Fireworks), prefix the id with the provider slash, e.g. "ollama/llama3.1:8b".${suggestionText}`,
|
||||
error: `Unknown model id "${trimmed}" for block "${blockType}". Read components/blocks/${blockType}.json (the model.options array) for valid ids; prefer entries with recommended: true and avoid deprecated: true. For user-configured models (Ollama, vLLM, LiteLLM, OpenRouter, Fireworks), prefix the id with the provider slash, e.g. "ollama/llama3.1:8b".${suggestionText}`,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -127,6 +127,8 @@ export const env = createEnv({
|
||||
OLLAMA_URL: z.string().url().optional(), // Ollama local LLM server URL
|
||||
VLLM_BASE_URL: z.string().url().optional(), // vLLM self-hosted base URL (OpenAI-compatible)
|
||||
VLLM_API_KEY: z.string().optional(), // Optional bearer token for vLLM
|
||||
LITELLM_BASE_URL: z.string().url().optional(), // LiteLLM proxy base URL (OpenAI-compatible)
|
||||
LITELLM_API_KEY: z.string().optional(), // Optional bearer token for LiteLLM
|
||||
FIREWORKS_API_KEY: z.string().optional(), // Optional Fireworks AI API key for model listing
|
||||
COHERE_API_KEY: z.string().min(1).optional(), // Cohere API key for reranker (rerank-v4.0-pro, rerank-v4.0-fast, rerank-v3.5)
|
||||
COHERE_API_KEY_1: z.string().min(1).optional(), // Primary Cohere API key for rotation
|
||||
|
||||
@@ -24,6 +24,7 @@ export type AttachmentProvider =
|
||||
| 'fireworks'
|
||||
| 'ollama'
|
||||
| 'vllm'
|
||||
| 'litellm'
|
||||
| 'xai'
|
||||
| 'deepseek'
|
||||
| 'cerebras'
|
||||
@@ -93,6 +94,7 @@ const PROVIDER_SUPPORTED_LABELS: Record<AttachmentProvider, string> = {
|
||||
fireworks: 'images through image_url message parts on vision models',
|
||||
ollama: 'images through image_url message parts on vision models',
|
||||
vllm: 'images through image_url message parts on multimodal models',
|
||||
litellm: 'images through image_url message parts on multimodal models',
|
||||
xai: 'images through image_url message parts on Grok vision models',
|
||||
deepseek: 'no file attachments in the current API adapter',
|
||||
cerebras: 'no file attachments in the current API adapter',
|
||||
@@ -109,6 +111,7 @@ export function getAttachmentProvider(providerId: ProviderId | string): Attachme
|
||||
if (providerId === 'fireworks') return 'fireworks'
|
||||
if (providerId === 'ollama') return 'ollama'
|
||||
if (providerId === 'vllm') return 'vllm'
|
||||
if (providerId === 'litellm') return 'litellm'
|
||||
if (providerId === 'xai') return 'xai'
|
||||
if (providerId === 'deepseek') return 'deepseek'
|
||||
if (providerId === 'cerebras') return 'cerebras'
|
||||
@@ -247,6 +250,7 @@ function isMimeTypeSupportedByProvider(
|
||||
case 'fireworks':
|
||||
case 'ollama':
|
||||
case 'vllm':
|
||||
case 'litellm':
|
||||
case 'xai':
|
||||
return isImageMimeType(mimeType)
|
||||
case 'deepseek':
|
||||
|
||||
@@ -0,0 +1,688 @@
|
||||
import { createLogger } from '@sim/logger'
|
||||
import { getErrorMessage, toError } from '@sim/utils/errors'
|
||||
import OpenAI from 'openai'
|
||||
import type { ChatCompletionCreateParamsStreaming } from 'openai/resources/chat/completions'
|
||||
import { env } from '@/lib/core/config/env'
|
||||
import type { StreamingExecution } from '@/executor/types'
|
||||
import { MAX_TOOL_ITERATIONS } from '@/providers'
|
||||
import { formatMessagesForProvider } from '@/providers/attachments'
|
||||
import { createReadableStreamFromLiteLLMStream } from '@/providers/litellm/utils'
|
||||
import { getProviderDefaultModel, getProviderModels } from '@/providers/models'
|
||||
import { enrichLastModelSegmentFromChatCompletions } from '@/providers/trace-enrichment'
|
||||
import type {
|
||||
Message,
|
||||
ProviderConfig,
|
||||
ProviderRequest,
|
||||
ProviderResponse,
|
||||
TimeSegment,
|
||||
} from '@/providers/types'
|
||||
import { ProviderError } from '@/providers/types'
|
||||
import {
|
||||
calculateCost,
|
||||
prepareToolExecution,
|
||||
prepareToolsWithUsageControl,
|
||||
sumToolCosts,
|
||||
trackForcedToolUsage,
|
||||
} from '@/providers/utils'
|
||||
import { useProvidersStore } from '@/stores/providers'
|
||||
import { executeTool } from '@/tools'
|
||||
|
||||
const logger = createLogger('LiteLLMProvider')
|
||||
const LITELLM_VERSION = '1.0.0'
|
||||
|
||||
export const litellmProvider: ProviderConfig = {
|
||||
id: 'litellm',
|
||||
name: 'LiteLLM',
|
||||
description: 'LiteLLM proxy with OpenAI-compatible API',
|
||||
version: LITELLM_VERSION,
|
||||
models: getProviderModels('litellm'),
|
||||
defaultModel: getProviderDefaultModel('litellm'),
|
||||
|
||||
async initialize() {
|
||||
if (typeof window !== 'undefined') {
|
||||
logger.info('Skipping LiteLLM initialization on client side to avoid CORS issues')
|
||||
return
|
||||
}
|
||||
|
||||
const baseUrl = (env.LITELLM_BASE_URL || '').replace(/\/$/, '')
|
||||
if (!baseUrl) {
|
||||
logger.info('LITELLM_BASE_URL not configured, skipping initialization')
|
||||
return
|
||||
}
|
||||
|
||||
try {
|
||||
const headers: Record<string, string> = {
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
|
||||
if (env.LITELLM_API_KEY) {
|
||||
headers.Authorization = `Bearer ${env.LITELLM_API_KEY}`
|
||||
}
|
||||
|
||||
const response = await fetch(`${baseUrl}/v1/models`, { headers })
|
||||
if (!response.ok) {
|
||||
await response.text().catch(() => {})
|
||||
useProvidersStore.getState().setProviderModels('litellm', [])
|
||||
logger.warn('LiteLLM service is not available. The provider will be disabled.')
|
||||
return
|
||||
}
|
||||
|
||||
const { vllmUpstreamResponseSchema } = await import('@/lib/api/contracts/providers')
|
||||
const data = vllmUpstreamResponseSchema.parse(await response.json())
|
||||
const models = data.data.map((model) => `litellm/${model.id}`)
|
||||
|
||||
this.models = models
|
||||
useProvidersStore.getState().setProviderModels('litellm', models)
|
||||
|
||||
logger.info(`Discovered ${models.length} LiteLLM model(s):`, { models })
|
||||
} catch (error) {
|
||||
logger.warn('LiteLLM model instantiation failed. The provider will be disabled.', {
|
||||
error: getErrorMessage(error, 'Unknown error'),
|
||||
})
|
||||
}
|
||||
},
|
||||
|
||||
executeRequest: async (
|
||||
request: ProviderRequest
|
||||
): Promise<ProviderResponse | StreamingExecution> => {
|
||||
logger.info('Preparing LiteLLM request', {
|
||||
model: request.model,
|
||||
hasSystemPrompt: !!request.systemPrompt,
|
||||
hasMessages: !!request.messages?.length,
|
||||
hasTools: !!request.tools?.length,
|
||||
toolCount: request.tools?.length || 0,
|
||||
hasResponseFormat: !!request.responseFormat,
|
||||
stream: !!request.stream,
|
||||
})
|
||||
|
||||
const baseUrl = (env.LITELLM_BASE_URL || '').replace(/\/$/, '')
|
||||
if (!baseUrl) {
|
||||
throw new Error('LITELLM_BASE_URL is required for LiteLLM provider')
|
||||
}
|
||||
|
||||
const apiKey = request.apiKey || env.LITELLM_API_KEY || 'empty'
|
||||
const litellm = new OpenAI({
|
||||
apiKey,
|
||||
baseURL: `${baseUrl}/v1`,
|
||||
})
|
||||
|
||||
const allMessages: Message[] = []
|
||||
|
||||
if (request.systemPrompt) {
|
||||
allMessages.push({
|
||||
role: 'system',
|
||||
content: request.systemPrompt,
|
||||
})
|
||||
}
|
||||
|
||||
if (request.context) {
|
||||
allMessages.push({
|
||||
role: 'user',
|
||||
content: request.context,
|
||||
})
|
||||
}
|
||||
|
||||
if (request.messages) {
|
||||
allMessages.push(...request.messages)
|
||||
}
|
||||
const formattedMessages = formatMessagesForProvider(allMessages, 'litellm') as Message[]
|
||||
|
||||
const tools = request.tools?.length
|
||||
? request.tools.map((tool) => ({
|
||||
type: 'function',
|
||||
function: {
|
||||
name: tool.id,
|
||||
description: tool.description,
|
||||
parameters: tool.parameters,
|
||||
},
|
||||
}))
|
||||
: undefined
|
||||
|
||||
const payload: any = {
|
||||
model: request.model.replace(/^litellm\//, ''),
|
||||
messages: formattedMessages,
|
||||
}
|
||||
|
||||
if (request.temperature !== undefined) payload.temperature = request.temperature
|
||||
if (request.maxTokens != null) payload.max_completion_tokens = request.maxTokens
|
||||
|
||||
if (request.responseFormat) {
|
||||
payload.response_format = {
|
||||
type: 'json_schema',
|
||||
json_schema: {
|
||||
name: request.responseFormat.name || 'response_schema',
|
||||
schema: request.responseFormat.schema || request.responseFormat,
|
||||
strict: request.responseFormat.strict !== false,
|
||||
},
|
||||
}
|
||||
|
||||
logger.info('Added JSON schema response format to LiteLLM request')
|
||||
}
|
||||
|
||||
let preparedTools: ReturnType<typeof prepareToolsWithUsageControl> | null = null
|
||||
let hasActiveTools = false
|
||||
|
||||
if (tools?.length) {
|
||||
preparedTools = prepareToolsWithUsageControl(tools, request.tools, logger, 'litellm')
|
||||
const { tools: filteredTools, toolChoice } = preparedTools
|
||||
|
||||
if (filteredTools?.length && toolChoice) {
|
||||
payload.tools = filteredTools
|
||||
payload.tool_choice = toolChoice
|
||||
hasActiveTools = true
|
||||
|
||||
logger.info('LiteLLM request configuration:', {
|
||||
toolCount: filteredTools.length,
|
||||
toolChoice:
|
||||
typeof toolChoice === 'string'
|
||||
? toolChoice
|
||||
: toolChoice.type === 'function'
|
||||
? `force:${toolChoice.function.name}`
|
||||
: 'unknown',
|
||||
model: payload.model,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
const providerStartTime = Date.now()
|
||||
const providerStartTimeISO = new Date(providerStartTime).toISOString()
|
||||
|
||||
try {
|
||||
if (request.stream && (!tools || tools.length === 0 || !hasActiveTools)) {
|
||||
logger.info('Using streaming response for LiteLLM request')
|
||||
|
||||
const streamingParams: ChatCompletionCreateParamsStreaming = {
|
||||
...payload,
|
||||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
}
|
||||
const streamResponse = await litellm.chat.completions.create(
|
||||
streamingParams,
|
||||
request.abortSignal ? { signal: request.abortSignal } : undefined
|
||||
)
|
||||
|
||||
const streamingResult = {
|
||||
stream: createReadableStreamFromLiteLLMStream(streamResponse, (content, usage) => {
|
||||
let cleanContent = content
|
||||
if (cleanContent && request.responseFormat) {
|
||||
cleanContent = cleanContent.replace(/```json\n?|\n?```/g, '').trim()
|
||||
}
|
||||
|
||||
streamingResult.execution.output.content = cleanContent
|
||||
streamingResult.execution.output.tokens = {
|
||||
input: usage.prompt_tokens,
|
||||
output: usage.completion_tokens,
|
||||
total: usage.total_tokens,
|
||||
}
|
||||
|
||||
const costResult = calculateCost(
|
||||
request.model,
|
||||
usage.prompt_tokens,
|
||||
usage.completion_tokens
|
||||
)
|
||||
streamingResult.execution.output.cost = {
|
||||
input: costResult.input,
|
||||
output: costResult.output,
|
||||
total: costResult.total,
|
||||
}
|
||||
|
||||
const streamEndTime = Date.now()
|
||||
const streamEndTimeISO = new Date(streamEndTime).toISOString()
|
||||
|
||||
if (streamingResult.execution.output.providerTiming) {
|
||||
streamingResult.execution.output.providerTiming.endTime = streamEndTimeISO
|
||||
streamingResult.execution.output.providerTiming.duration =
|
||||
streamEndTime - providerStartTime
|
||||
|
||||
if (streamingResult.execution.output.providerTiming.timeSegments?.[0]) {
|
||||
streamingResult.execution.output.providerTiming.timeSegments[0].endTime =
|
||||
streamEndTime
|
||||
streamingResult.execution.output.providerTiming.timeSegments[0].duration =
|
||||
streamEndTime - providerStartTime
|
||||
}
|
||||
}
|
||||
}),
|
||||
execution: {
|
||||
success: true,
|
||||
output: {
|
||||
content: '',
|
||||
model: request.model,
|
||||
tokens: { input: 0, output: 0, total: 0 },
|
||||
toolCalls: undefined,
|
||||
providerTiming: {
|
||||
startTime: providerStartTimeISO,
|
||||
endTime: new Date().toISOString(),
|
||||
duration: Date.now() - providerStartTime,
|
||||
timeSegments: [
|
||||
{
|
||||
type: 'model',
|
||||
name: request.model,
|
||||
startTime: providerStartTime,
|
||||
endTime: Date.now(),
|
||||
duration: Date.now() - providerStartTime,
|
||||
},
|
||||
],
|
||||
},
|
||||
cost: { input: 0, output: 0, total: 0 },
|
||||
},
|
||||
logs: [],
|
||||
metadata: {
|
||||
startTime: providerStartTimeISO,
|
||||
endTime: new Date().toISOString(),
|
||||
duration: Date.now() - providerStartTime,
|
||||
},
|
||||
},
|
||||
} as StreamingExecution
|
||||
|
||||
return streamingResult as StreamingExecution
|
||||
}
|
||||
|
||||
const initialCallTime = Date.now()
|
||||
|
||||
const originalToolChoice = payload.tool_choice
|
||||
|
||||
const forcedTools = preparedTools?.forcedTools || []
|
||||
let usedForcedTools: string[] = []
|
||||
|
||||
const checkForForcedToolUsage = (
|
||||
response: any,
|
||||
toolChoice: string | { type: string; function?: { name: string }; name?: string; any?: any }
|
||||
) => {
|
||||
if (typeof toolChoice === 'object' && response.choices[0]?.message?.tool_calls) {
|
||||
const toolCallsResponse = response.choices[0].message.tool_calls
|
||||
const result = trackForcedToolUsage(
|
||||
toolCallsResponse,
|
||||
toolChoice,
|
||||
logger,
|
||||
'litellm',
|
||||
forcedTools,
|
||||
usedForcedTools
|
||||
)
|
||||
hasUsedForcedTool = result.hasUsedForcedTool
|
||||
usedForcedTools = result.usedForcedTools
|
||||
}
|
||||
}
|
||||
|
||||
let currentResponse = await litellm.chat.completions.create(
|
||||
payload,
|
||||
request.abortSignal ? { signal: request.abortSignal } : undefined
|
||||
)
|
||||
const firstResponseTime = Date.now() - initialCallTime
|
||||
|
||||
let content = currentResponse.choices[0]?.message?.content || ''
|
||||
|
||||
if (content && request.responseFormat) {
|
||||
content = content.replace(/```json\n?|\n?```/g, '').trim()
|
||||
}
|
||||
|
||||
const tokens = {
|
||||
input: currentResponse.usage?.prompt_tokens || 0,
|
||||
output: currentResponse.usage?.completion_tokens || 0,
|
||||
total: currentResponse.usage?.total_tokens || 0,
|
||||
}
|
||||
const toolCalls = []
|
||||
const toolResults: Record<string, unknown>[] = []
|
||||
const currentMessages = [...formattedMessages]
|
||||
let iterationCount = 0
|
||||
|
||||
let modelTime = firstResponseTime
|
||||
let toolsTime = 0
|
||||
|
||||
let hasUsedForcedTool = false
|
||||
|
||||
const timeSegments: TimeSegment[] = [
|
||||
{
|
||||
type: 'model',
|
||||
name: request.model,
|
||||
startTime: initialCallTime,
|
||||
endTime: initialCallTime + firstResponseTime,
|
||||
duration: firstResponseTime,
|
||||
},
|
||||
]
|
||||
|
||||
checkForForcedToolUsage(currentResponse, originalToolChoice)
|
||||
|
||||
while (iterationCount < MAX_TOOL_ITERATIONS) {
|
||||
if (currentResponse.choices[0]?.message?.content) {
|
||||
content = currentResponse.choices[0].message.content
|
||||
if (request.responseFormat) {
|
||||
content = content.replace(/```json\n?|\n?```/g, '').trim()
|
||||
}
|
||||
}
|
||||
|
||||
const toolCallsInResponse = currentResponse.choices[0]?.message?.tool_calls
|
||||
|
||||
enrichLastModelSegmentFromChatCompletions(
|
||||
timeSegments,
|
||||
currentResponse,
|
||||
toolCallsInResponse,
|
||||
{ model: request.model, provider: 'litellm' }
|
||||
)
|
||||
|
||||
if (!toolCallsInResponse || toolCallsInResponse.length === 0) {
|
||||
break
|
||||
}
|
||||
|
||||
logger.info(
|
||||
`Processing ${toolCallsInResponse.length} tool calls (iteration ${iterationCount + 1}/${MAX_TOOL_ITERATIONS})`
|
||||
)
|
||||
|
||||
const toolsStartTime = Date.now()
|
||||
|
||||
const toolExecutionPromises = toolCallsInResponse.map(async (toolCall) => {
|
||||
const toolCallStartTime = Date.now()
|
||||
const toolName = toolCall.function.name
|
||||
|
||||
try {
|
||||
const toolArgs = JSON.parse(toolCall.function.arguments)
|
||||
const tool = request.tools?.find((t) => t.id === toolName)
|
||||
|
||||
if (!tool) return null
|
||||
|
||||
const { toolParams, executionParams } = prepareToolExecution(tool, toolArgs, request)
|
||||
const result = await executeTool(toolName, executionParams, {
|
||||
signal: request.abortSignal,
|
||||
})
|
||||
const toolCallEndTime = Date.now()
|
||||
|
||||
return {
|
||||
toolCall,
|
||||
toolName,
|
||||
toolParams,
|
||||
result,
|
||||
startTime: toolCallStartTime,
|
||||
endTime: toolCallEndTime,
|
||||
duration: toolCallEndTime - toolCallStartTime,
|
||||
}
|
||||
} catch (error) {
|
||||
const toolCallEndTime = Date.now()
|
||||
logger.error('Error processing tool call:', { error, toolName })
|
||||
|
||||
return {
|
||||
toolCall,
|
||||
toolName,
|
||||
toolParams: {},
|
||||
result: {
|
||||
success: false,
|
||||
output: undefined,
|
||||
error: getErrorMessage(error, 'Tool execution failed'),
|
||||
},
|
||||
startTime: toolCallStartTime,
|
||||
endTime: toolCallEndTime,
|
||||
duration: toolCallEndTime - toolCallStartTime,
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const executionResults = await Promise.allSettled(toolExecutionPromises)
|
||||
|
||||
currentMessages.push({
|
||||
role: 'assistant',
|
||||
content: null,
|
||||
tool_calls: toolCallsInResponse.map((tc) => ({
|
||||
id: tc.id,
|
||||
type: 'function',
|
||||
function: {
|
||||
name: tc.function.name,
|
||||
arguments: tc.function.arguments,
|
||||
},
|
||||
})),
|
||||
})
|
||||
|
||||
for (const settledResult of executionResults) {
|
||||
if (settledResult.status === 'rejected' || !settledResult.value) continue
|
||||
|
||||
const { toolCall, toolName, toolParams, result, startTime, endTime, duration } =
|
||||
settledResult.value
|
||||
|
||||
timeSegments.push({
|
||||
type: 'tool',
|
||||
name: toolName,
|
||||
startTime: startTime,
|
||||
endTime: endTime,
|
||||
duration: duration,
|
||||
toolCallId: toolCall.id,
|
||||
})
|
||||
|
||||
let resultContent: any
|
||||
if (result.success && result.output) {
|
||||
toolResults.push(result.output)
|
||||
resultContent = result.output
|
||||
} else {
|
||||
resultContent = {
|
||||
error: true,
|
||||
message: result.error || 'Tool execution failed',
|
||||
tool: toolName,
|
||||
}
|
||||
}
|
||||
|
||||
toolCalls.push({
|
||||
name: toolName,
|
||||
arguments: toolParams,
|
||||
startTime: new Date(startTime).toISOString(),
|
||||
endTime: new Date(endTime).toISOString(),
|
||||
duration: duration,
|
||||
result: resultContent,
|
||||
success: result.success,
|
||||
})
|
||||
|
||||
currentMessages.push({
|
||||
role: 'tool',
|
||||
tool_call_id: toolCall.id,
|
||||
content: JSON.stringify(resultContent),
|
||||
})
|
||||
}
|
||||
|
||||
const thisToolsTime = Date.now() - toolsStartTime
|
||||
toolsTime += thisToolsTime
|
||||
|
||||
const nextPayload = {
|
||||
...payload,
|
||||
messages: currentMessages,
|
||||
}
|
||||
|
||||
if (typeof originalToolChoice === 'object' && hasUsedForcedTool && forcedTools.length > 0) {
|
||||
const remainingTools = forcedTools.filter((tool) => !usedForcedTools.includes(tool))
|
||||
|
||||
if (remainingTools.length > 0) {
|
||||
nextPayload.tool_choice = {
|
||||
type: 'function',
|
||||
function: { name: remainingTools[0] },
|
||||
}
|
||||
logger.info(`Forcing next tool: ${remainingTools[0]}`)
|
||||
} else {
|
||||
nextPayload.tool_choice = 'auto'
|
||||
logger.info('All forced tools have been used, switching to auto tool_choice')
|
||||
}
|
||||
}
|
||||
|
||||
const nextModelStartTime = Date.now()
|
||||
|
||||
currentResponse = await litellm.chat.completions.create(
|
||||
nextPayload,
|
||||
request.abortSignal ? { signal: request.abortSignal } : undefined
|
||||
)
|
||||
|
||||
checkForForcedToolUsage(currentResponse, nextPayload.tool_choice)
|
||||
|
||||
const nextModelEndTime = Date.now()
|
||||
const thisModelTime = nextModelEndTime - nextModelStartTime
|
||||
|
||||
timeSegments.push({
|
||||
type: 'model',
|
||||
name: request.model,
|
||||
startTime: nextModelStartTime,
|
||||
endTime: nextModelEndTime,
|
||||
duration: thisModelTime,
|
||||
})
|
||||
|
||||
modelTime += thisModelTime
|
||||
|
||||
if (currentResponse.choices[0]?.message?.content) {
|
||||
content = currentResponse.choices[0].message.content
|
||||
if (request.responseFormat) {
|
||||
content = content.replace(/```json\n?|\n?```/g, '').trim()
|
||||
}
|
||||
}
|
||||
|
||||
if (currentResponse.usage) {
|
||||
tokens.input += currentResponse.usage.prompt_tokens || 0
|
||||
tokens.output += currentResponse.usage.completion_tokens || 0
|
||||
tokens.total += currentResponse.usage.total_tokens || 0
|
||||
}
|
||||
|
||||
iterationCount++
|
||||
}
|
||||
|
||||
if (iterationCount === MAX_TOOL_ITERATIONS) {
|
||||
enrichLastModelSegmentFromChatCompletions(
|
||||
timeSegments,
|
||||
currentResponse,
|
||||
currentResponse.choices[0]?.message?.tool_calls,
|
||||
{ model: request.model, provider: 'litellm' }
|
||||
)
|
||||
}
|
||||
|
||||
if (request.stream) {
|
||||
logger.info('Using streaming for final response after tool processing')
|
||||
|
||||
const accumulatedCost = calculateCost(request.model, tokens.input, tokens.output)
|
||||
|
||||
const streamingParams: ChatCompletionCreateParamsStreaming = {
|
||||
...payload,
|
||||
messages: currentMessages,
|
||||
tool_choice: 'auto',
|
||||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
}
|
||||
const streamResponse = await litellm.chat.completions.create(
|
||||
streamingParams,
|
||||
request.abortSignal ? { signal: request.abortSignal } : undefined
|
||||
)
|
||||
|
||||
const streamingResult = {
|
||||
stream: createReadableStreamFromLiteLLMStream(streamResponse, (content, usage) => {
|
||||
let cleanContent = content
|
||||
if (cleanContent && request.responseFormat) {
|
||||
cleanContent = cleanContent.replace(/```json\n?|\n?```/g, '').trim()
|
||||
}
|
||||
|
||||
streamingResult.execution.output.content = cleanContent
|
||||
streamingResult.execution.output.tokens = {
|
||||
input: tokens.input + usage.prompt_tokens,
|
||||
output: tokens.output + usage.completion_tokens,
|
||||
total: tokens.total + usage.total_tokens,
|
||||
}
|
||||
|
||||
const streamCost = calculateCost(
|
||||
request.model,
|
||||
usage.prompt_tokens,
|
||||
usage.completion_tokens
|
||||
)
|
||||
const tc = sumToolCosts(toolResults)
|
||||
streamingResult.execution.output.cost = {
|
||||
input: accumulatedCost.input + streamCost.input,
|
||||
output: accumulatedCost.output + streamCost.output,
|
||||
toolCost: tc || undefined,
|
||||
total: accumulatedCost.total + streamCost.total + tc,
|
||||
}
|
||||
}),
|
||||
execution: {
|
||||
success: true,
|
||||
output: {
|
||||
content: '',
|
||||
model: request.model,
|
||||
tokens: {
|
||||
input: tokens.input,
|
||||
output: tokens.output,
|
||||
total: tokens.total,
|
||||
},
|
||||
toolCalls:
|
||||
toolCalls.length > 0
|
||||
? {
|
||||
list: toolCalls,
|
||||
count: toolCalls.length,
|
||||
}
|
||||
: undefined,
|
||||
providerTiming: {
|
||||
startTime: providerStartTimeISO,
|
||||
endTime: new Date().toISOString(),
|
||||
duration: Date.now() - providerStartTime,
|
||||
modelTime: modelTime,
|
||||
toolsTime: toolsTime,
|
||||
firstResponseTime: firstResponseTime,
|
||||
iterations: iterationCount + 1,
|
||||
timeSegments: timeSegments,
|
||||
},
|
||||
cost: {
|
||||
input: accumulatedCost.input,
|
||||
output: accumulatedCost.output,
|
||||
total: accumulatedCost.total,
|
||||
},
|
||||
},
|
||||
logs: [],
|
||||
metadata: {
|
||||
startTime: providerStartTimeISO,
|
||||
endTime: new Date().toISOString(),
|
||||
duration: Date.now() - providerStartTime,
|
||||
},
|
||||
},
|
||||
} as StreamingExecution
|
||||
|
||||
return streamingResult as StreamingExecution
|
||||
}
|
||||
|
||||
const providerEndTime = Date.now()
|
||||
const providerEndTimeISO = new Date(providerEndTime).toISOString()
|
||||
const totalDuration = providerEndTime - providerStartTime
|
||||
|
||||
return {
|
||||
content,
|
||||
model: request.model,
|
||||
tokens,
|
||||
toolCalls: toolCalls.length > 0 ? toolCalls : undefined,
|
||||
toolResults: toolResults.length > 0 ? toolResults : undefined,
|
||||
timing: {
|
||||
startTime: providerStartTimeISO,
|
||||
endTime: providerEndTimeISO,
|
||||
duration: totalDuration,
|
||||
modelTime: modelTime,
|
||||
toolsTime: toolsTime,
|
||||
firstResponseTime: firstResponseTime,
|
||||
iterations: iterationCount + 1,
|
||||
timeSegments: timeSegments,
|
||||
},
|
||||
}
|
||||
} catch (error) {
|
||||
const providerEndTime = Date.now()
|
||||
const providerEndTimeISO = new Date(providerEndTime).toISOString()
|
||||
const totalDuration = providerEndTime - providerStartTime
|
||||
|
||||
let errorMessage = toError(error).message
|
||||
let errorType: string | undefined
|
||||
let errorCode: number | undefined
|
||||
|
||||
if (error && typeof error === 'object' && 'error' in error) {
|
||||
const litellmError = error.error as any
|
||||
if (litellmError && typeof litellmError === 'object') {
|
||||
errorMessage = litellmError.message || errorMessage
|
||||
errorType = litellmError.type
|
||||
errorCode = litellmError.code
|
||||
}
|
||||
}
|
||||
|
||||
logger.error('Error in LiteLLM request:', {
|
||||
error: errorMessage,
|
||||
errorType,
|
||||
errorCode,
|
||||
duration: totalDuration,
|
||||
})
|
||||
|
||||
throw new ProviderError(errorMessage, {
|
||||
startTime: providerStartTimeISO,
|
||||
endTime: providerEndTimeISO,
|
||||
duration: totalDuration,
|
||||
})
|
||||
}
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
import type { ChatCompletionChunk } from 'openai/resources/chat/completions'
|
||||
import type { CompletionUsage } from 'openai/resources/completions'
|
||||
import { createOpenAICompatibleStream } from '@/providers/utils'
|
||||
|
||||
/**
|
||||
* Creates a ReadableStream from a LiteLLM streaming response.
|
||||
* Uses the shared OpenAI-compatible streaming utility.
|
||||
*/
|
||||
export function createReadableStreamFromLiteLLMStream(
|
||||
litellmStream: AsyncIterable<ChatCompletionChunk>,
|
||||
onComplete?: (content: string, usage: CompletionUsage) => void
|
||||
): ReadableStream<Uint8Array> {
|
||||
return createOpenAICompatibleStream(litellmStream, 'LiteLLM', onComplete)
|
||||
}
|
||||
@@ -17,6 +17,7 @@ import {
|
||||
FireworksIcon,
|
||||
GeminiIcon,
|
||||
GroqIcon,
|
||||
LitellmIcon,
|
||||
MistralIcon,
|
||||
OllamaIcon,
|
||||
OpenAIIcon,
|
||||
@@ -125,6 +126,20 @@ export const PROVIDER_DEFINITIONS: Record<string, ProviderDefinition> = {
|
||||
},
|
||||
models: [],
|
||||
},
|
||||
litellm: {
|
||||
id: 'litellm',
|
||||
name: 'LiteLLM',
|
||||
icon: LitellmIcon,
|
||||
color: '#040229',
|
||||
description: 'LiteLLM proxy with an OpenAI-compatible API',
|
||||
defaultModel: '',
|
||||
modelPatterns: [/^litellm\//],
|
||||
capabilities: {
|
||||
temperature: { min: 0, max: 2 },
|
||||
toolUsageControl: true,
|
||||
},
|
||||
models: [],
|
||||
},
|
||||
openai: {
|
||||
id: 'openai',
|
||||
name: 'OpenAI',
|
||||
@@ -2803,7 +2818,13 @@ export function getProviderModels(providerId: string): string[] {
|
||||
return PROVIDER_DEFINITIONS[providerId]?.models.map((m) => m.id) || []
|
||||
}
|
||||
|
||||
export const DYNAMIC_MODEL_PROVIDERS = ['ollama', 'vllm', 'openrouter', 'fireworks'] as const
|
||||
export const DYNAMIC_MODEL_PROVIDERS = [
|
||||
'ollama',
|
||||
'vllm',
|
||||
'litellm',
|
||||
'openrouter',
|
||||
'fireworks',
|
||||
] as const
|
||||
|
||||
function getAllStaticModelIds(): string[] {
|
||||
const ids: string[] = []
|
||||
@@ -2857,7 +2878,7 @@ export function suggestModelIdsForUnknownModel(_modelId: string, limit = 5): str
|
||||
|
||||
export function getBaseModelProviders(): Record<string, ProviderId> {
|
||||
return Object.entries(PROVIDER_DEFINITIONS)
|
||||
.filter(([providerId]) => !['ollama', 'vllm', 'openrouter'].includes(providerId))
|
||||
.filter(([providerId]) => !['ollama', 'vllm', 'litellm', 'openrouter'].includes(providerId))
|
||||
.reduce(
|
||||
(map, [providerId, provider]) => {
|
||||
provider.models.forEach((model) => {
|
||||
@@ -3034,6 +3055,18 @@ export function updateVLLMModels(models: string[]): void {
|
||||
}))
|
||||
}
|
||||
|
||||
export function updateLiteLLMModels(models: string[]): void {
|
||||
PROVIDER_DEFINITIONS.litellm.models = models.map((modelId) => ({
|
||||
id: modelId,
|
||||
pricing: {
|
||||
input: 0,
|
||||
output: 0,
|
||||
updatedAt: new Date().toISOString().split('T')[0],
|
||||
},
|
||||
capabilities: {},
|
||||
}))
|
||||
}
|
||||
|
||||
export function updateFireworksModels(models: string[]): void {
|
||||
PROVIDER_DEFINITIONS.fireworks.models = models.map((modelId) => ({
|
||||
id: modelId,
|
||||
|
||||
@@ -9,6 +9,7 @@ import { deepseekProvider } from '@/providers/deepseek'
|
||||
import { fireworksProvider } from '@/providers/fireworks'
|
||||
import { googleProvider } from '@/providers/google'
|
||||
import { groqProvider } from '@/providers/groq'
|
||||
import { litellmProvider } from '@/providers/litellm'
|
||||
import { mistralProvider } from '@/providers/mistral'
|
||||
import { ollamaProvider } from '@/providers/ollama'
|
||||
import { openaiProvider } from '@/providers/openai'
|
||||
@@ -31,6 +32,7 @@ const providerRegistry: Record<ProviderId, ProviderConfig> = {
|
||||
cerebras: cerebrasProvider,
|
||||
groq: groqProvider,
|
||||
vllm: vllmProvider,
|
||||
litellm: litellmProvider,
|
||||
mistral: mistralProvider,
|
||||
'azure-openai': azureOpenAIProvider,
|
||||
openrouter: openRouterProvider,
|
||||
|
||||
@@ -16,6 +16,7 @@ export type ProviderId =
|
||||
| 'openrouter'
|
||||
| 'fireworks'
|
||||
| 'vllm'
|
||||
| 'litellm'
|
||||
| 'bedrock'
|
||||
|
||||
export interface ModelPricing {
|
||||
|
||||
@@ -168,6 +168,19 @@ describe('getApiKey', () => {
|
||||
expect(key2).toBe('user-key')
|
||||
}
|
||||
)
|
||||
|
||||
it.concurrent(
|
||||
'should return empty or user-provided key for litellm provider without requiring API key',
|
||||
() => {
|
||||
isHostedSpy.mockReturnValue(false)
|
||||
|
||||
const key = getApiKey('litellm', 'litellm/anthropic/claude-sonnet-4-6')
|
||||
expect(key).toBe('empty')
|
||||
|
||||
const key2 = getApiKey('litellm', 'litellm/openai/gpt-4', 'user-key')
|
||||
expect(key2).toBe('user-key')
|
||||
}
|
||||
)
|
||||
})
|
||||
|
||||
describe('Model Capabilities', () => {
|
||||
|
||||
@@ -132,6 +132,7 @@ function buildProviderMetadata(providerId: ProviderId): ProviderMetadata {
|
||||
export const providers: Record<ProviderId, ProviderMetadata> = {
|
||||
ollama: buildProviderMetadata('ollama'),
|
||||
vllm: buildProviderMetadata('vllm'),
|
||||
litellm: buildProviderMetadata('litellm'),
|
||||
openai: {
|
||||
...buildProviderMetadata('openai'),
|
||||
computerUseModels: ['computer-use-preview'],
|
||||
@@ -167,6 +168,12 @@ export function updateVLLMProviderModels(models: string[]): void {
|
||||
providers.vllm.models = getProviderModelsFromDefinitions('vllm')
|
||||
}
|
||||
|
||||
export function updateLiteLLMProviderModels(models: string[]): void {
|
||||
const { updateLiteLLMModels } = require('@/providers/models')
|
||||
updateLiteLLMModels(models)
|
||||
providers.litellm.models = getProviderModelsFromDefinitions('litellm')
|
||||
}
|
||||
|
||||
export async function updateOpenRouterProviderModels(models: string[]): Promise<void> {
|
||||
const { updateOpenRouterModels } = await import('@/providers/models')
|
||||
updateOpenRouterModels(models)
|
||||
@@ -185,6 +192,7 @@ export function getBaseModelProviders(): Record<string, ProviderId> {
|
||||
([providerId]) =>
|
||||
providerId !== 'ollama' &&
|
||||
providerId !== 'vllm' &&
|
||||
providerId !== 'litellm' &&
|
||||
providerId !== 'openrouter' &&
|
||||
providerId !== 'fireworks'
|
||||
)
|
||||
@@ -744,6 +752,12 @@ export function getApiKey(provider: string, model: string, userProvidedKey?: str
|
||||
return userProvidedKey || 'empty'
|
||||
}
|
||||
|
||||
const isLitellmModel =
|
||||
provider === 'litellm' || useProvidersStore.getState().providers.litellm.models.includes(model)
|
||||
if (isLitellmModel) {
|
||||
return userProvidedKey || 'empty'
|
||||
}
|
||||
|
||||
// Bedrock uses its own credentials (bedrockAccessKeyId/bedrockSecretKey), not apiKey
|
||||
const isBedrockModel = provider === 'bedrock' || model.startsWith('bedrock/')
|
||||
if (isBedrockModel) {
|
||||
|
||||
@@ -9,6 +9,7 @@ export const useProvidersStore = create<ProvidersStore>((set, get) => ({
|
||||
base: { models: [], isLoading: false },
|
||||
ollama: { models: [], isLoading: false },
|
||||
vllm: { models: [], isLoading: false },
|
||||
litellm: { models: [], isLoading: false },
|
||||
openrouter: { models: [], isLoading: false },
|
||||
fireworks: { models: [], isLoading: false },
|
||||
},
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
export type ProviderName = 'ollama' | 'vllm' | 'openrouter' | 'fireworks' | 'base'
|
||||
export type ProviderName = 'ollama' | 'vllm' | 'litellm' | 'openrouter' | 'fireworks' | 'base'
|
||||
|
||||
export interface OpenRouterModelInfo {
|
||||
id: string
|
||||
|
||||
@@ -9,8 +9,8 @@ const QUERY_HOOKS_DIR = path.join(ROOT, 'apps/sim/hooks/queries')
|
||||
const SELECTOR_HOOKS_DIR = path.join(ROOT, 'apps/sim/hooks/selectors')
|
||||
|
||||
const BASELINE = {
|
||||
totalRoutes: 755,
|
||||
zodRoutes: 755,
|
||||
totalRoutes: 756,
|
||||
zodRoutes: 756,
|
||||
nonZodRoutes: 0,
|
||||
} as const
|
||||
|
||||
|
||||
Reference in New Issue
Block a user