| |
| |
| |
| |
| |
|
|
| import type { GenerateContentConfig } from '@google/genai'; |
| import type { ModelPolicy } from '../availability/modelPolicy.js'; |
| import { |
| getDisplayString, |
| PREVIEW_GEMINI_3_1_MODEL, |
| isProModel, |
| getAutoModelDescription, |
| } from '../config/models.js'; |
|
|
| |
| |
| export interface ModelConfigKey { |
| model: string; |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| overrideScope?: string; |
|
|
| |
| |
| |
| isRetry?: boolean; |
|
|
| |
| |
| isChatModel?: boolean; |
|
|
| |
| lastStreamError?: unknown; |
| } |
|
|
| export interface ModelConfig { |
| model?: string; |
| generateContentConfig?: GenerateContentConfig; |
| } |
|
|
| export interface ModelConfigOverride { |
| match: { |
| model?: string; |
| overrideScope?: string; |
| isRetry?: boolean; |
| }; |
| modelConfig: ModelConfig; |
| } |
|
|
| export interface ModelConfigAlias { |
| extends?: string; |
| modelConfig: ModelConfig; |
| } |
|
|
| |
| |
| |
| export interface ModelDefinition { |
| displayName?: string; |
| tier?: string; |
| family?: string; |
| isPreview?: boolean; |
| |
| isVisible?: boolean; |
| |
| dialogDescription?: string; |
| features?: { |
| |
| thinking?: boolean; |
| |
| |
| multimodalToolUse?: boolean; |
| }; |
| } |
|
|
| |
| |
| export interface ModelResolution { |
| |
| default: string; |
| |
| contexts?: Array<{ |
| |
| condition: ResolutionCondition; |
| |
| target: string; |
| }>; |
| } |
|
|
| |
| export interface ResolutionContext { |
| useGemini3_1?: boolean; |
| useGemini3_1FlashLite?: boolean; |
| useGemini3_5Flash?: boolean; |
| useCustomTools?: boolean; |
| hasAccessToPreview?: boolean; |
| hasAccessToProModel?: boolean; |
| requestedModel?: string; |
| } |
|
|
| |
| export interface ResolutionCondition { |
| useGemini3_1?: boolean; |
| useGemini3_1FlashLite?: boolean; |
| useGemini3_5Flash?: boolean; |
| useCustomTools?: boolean; |
| hasAccessToPreview?: boolean; |
| |
| requestedModels?: string[]; |
| } |
|
|
| export interface ModelConfigServiceConfig { |
| aliases?: Record<string, ModelConfigAlias>; |
| customAliases?: Record<string, ModelConfigAlias>; |
| overrides?: ModelConfigOverride[]; |
| customOverrides?: ModelConfigOverride[]; |
| modelDefinitions?: Record<string, ModelDefinition>; |
| modelIdResolutions?: Record<string, ModelResolution>; |
| classifierIdResolutions?: Record<string, ModelResolution>; |
| modelChains?: Record<string, ModelPolicy[]>; |
| } |
|
|
| const MAX_ALIAS_CHAIN_DEPTH = 100; |
|
|
| export type ResolvedModelConfig = _ResolvedModelConfig & { |
| readonly _brand: unique symbol; |
| }; |
|
|
| export interface _ResolvedModelConfig { |
| model: string; |
| generateContentConfig: GenerateContentConfig; |
| } |
|
|
| export class ModelConfigService { |
| private readonly runtimeAliases: Record<string, ModelConfigAlias> = {}; |
| private readonly runtimeOverrides: ModelConfigOverride[] = []; |
|
|
| |
| constructor(private readonly config: ModelConfigServiceConfig) {} |
|
|
| |
| |
| |
| |
| getAvailableModelOptions(context: ResolutionContext): Array<{ |
| modelId: string; |
| name: string; |
| description: string; |
| tier: string; |
| }> { |
| const definitions = this.config.modelDefinitions ?? {}; |
| const shouldShowPreviewModels = context.hasAccessToPreview ?? false; |
| const useGemini31 = context.useGemini3_1 ?? false; |
| const useGemini3_5Flash = context.useGemini3_5Flash ?? false; |
|
|
| const mainOptions = Object.entries(definitions) |
| .filter(([_, m]) => { |
| if (m.isVisible !== true) return false; |
| if (m.isPreview && !shouldShowPreviewModels) return false; |
| if (m.tier !== 'auto') return false; |
| return true; |
| }) |
| .map(([id, m]) => { |
| let description = m.dialogDescription ?? ''; |
| if (id === 'auto') { |
| description = getAutoModelDescription( |
| shouldShowPreviewModels, |
| useGemini31, |
| useGemini3_5Flash, |
| ); |
| } else if (id === 'auto-gemini-3' && useGemini31) { |
| description = description.replace('gemini-3-pro', 'gemini-3.1-pro'); |
| } |
|
|
| return { |
| modelId: id, |
| name: m.displayName ?? getDisplayString(id), |
| description, |
| tier: m.tier ?? 'auto', |
| }; |
| }); |
|
|
| const manualOptions = Object.entries(definitions) |
| .filter(([id, m]) => { |
| if (m.isVisible !== true) return false; |
| if (m.isPreview && !shouldShowPreviewModels) return false; |
| if (m.tier === 'auto') return false; |
| if (context.hasAccessToProModel === false && isProModel(id)) |
| return false; |
| if (id === PREVIEW_GEMINI_3_1_MODEL && !useGemini31) return false; |
| return true; |
| }) |
| .map(([id, m]) => { |
| const resolvedId = this.resolveModelId(id, context); |
| const titleId = this.resolveModelId(id, { |
| useGemini3_1: useGemini31, |
| }); |
| return { |
| modelId: resolvedId, |
| name: m.displayName ?? getDisplayString(titleId), |
| description: m.dialogDescription ?? '', |
| tier: m.tier ?? 'custom', |
| }; |
| }); |
|
|
| |
| const seen = new Set<string>(); |
| const uniqueManualOptions = manualOptions.filter((option) => { |
| if (seen.has(option.modelId)) return false; |
| seen.add(option.modelId); |
| return true; |
| }); |
|
|
| return [...mainOptions, ...uniqueManualOptions]; |
| } |
|
|
| getModelDefinition(modelId: string): ModelDefinition | undefined { |
| const definition = this.config.modelDefinitions?.[modelId]; |
| if (definition) { |
| return definition; |
| } |
|
|
| |
| if (!modelId.startsWith('gemini-')) { |
| return { |
| tier: 'custom', |
| family: 'custom', |
| features: {}, |
| }; |
| } |
|
|
| return undefined; |
| } |
|
|
| getModelDefinitions(): Record<string, ModelDefinition> { |
| return this.config.modelDefinitions ?? {}; |
| } |
|
|
| private matches( |
| condition: ResolutionCondition, |
| context: ResolutionContext, |
| ): boolean { |
| return Object.entries(condition).every(([key, value]) => { |
| if (value === undefined) return true; |
|
|
| switch (key) { |
| case 'useGemini3_1': |
| return value === context.useGemini3_1; |
| case 'useGemini3_1FlashLite': |
| return value === context.useGemini3_1FlashLite; |
| case 'useGemini3_5Flash': |
| return value === context.useGemini3_5Flash; |
| case 'useCustomTools': |
| return value === context.useCustomTools; |
| case 'hasAccessToPreview': |
| return value === context.hasAccessToPreview; |
| case 'requestedModels': |
| return ( |
| Array.isArray(value) && |
| !!context.requestedModel && |
| value.includes(context.requestedModel) |
| ); |
| default: |
| return false; |
| } |
| }); |
| } |
|
|
| |
| resolveModelId( |
| requestedName: string, |
| context: ResolutionContext = {}, |
| ): string { |
| const resolution = this.config.modelIdResolutions?.[requestedName]; |
| if (!resolution) { |
| return requestedName; |
| } |
|
|
| for (const ctx of resolution.contexts ?? []) { |
| if (this.matches(ctx.condition, context)) { |
| return ctx.target; |
| } |
| } |
|
|
| return resolution.default; |
| } |
|
|
| |
| resolveClassifierModelId( |
| tier: string, |
| requestedModel: string, |
| context: ResolutionContext = {}, |
| ): string { |
| const resolution = this.config.classifierIdResolutions?.[tier]; |
| const fullContext: ResolutionContext = { ...context, requestedModel }; |
|
|
| if (!resolution) { |
| |
| return this.resolveModelId(tier, fullContext); |
| } |
|
|
| for (const ctx of resolution.contexts ?? []) { |
| if (this.matches(ctx.condition, fullContext)) { |
| return ctx.target; |
| } |
| } |
|
|
| return resolution.default; |
| } |
|
|
| getModelChain(chainName: string): ModelPolicy[] | undefined { |
| return this.config.modelChains?.[chainName]; |
| } |
|
|
| |
| |
| |
| |
| resolveChain( |
| chainName: string, |
| context: ResolutionContext = {}, |
| ): ModelPolicy[] | undefined { |
| const template = this.config.modelChains?.[chainName]; |
| if (!template) { |
| return undefined; |
| } |
| |
| return template.map((policy) => ({ |
| ...policy, |
| model: this.resolveModelId(policy.model, context), |
| })); |
| } |
|
|
| registerRuntimeModelConfig(aliasName: string, alias: ModelConfigAlias): void { |
| this.runtimeAliases[aliasName] = alias; |
| } |
|
|
| registerRuntimeModelOverride(override: ModelConfigOverride): void { |
| this.runtimeOverrides.push(override); |
| } |
|
|
| clearRuntimeOverrides(): void { |
| this.runtimeOverrides.length = 0; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| private internalGetResolvedConfig(context: ModelConfigKey): { |
| model: string | undefined; |
| generateContentConfig: GenerateContentConfig; |
| } { |
| const { |
| aliases = {}, |
| customAliases = {}, |
| overrides = [], |
| customOverrides = [], |
| } = this.config || {}; |
| const allAliases = { |
| ...aliases, |
| ...customAliases, |
| ...this.runtimeAliases, |
| }; |
|
|
| const { aliasChain, baseModel, resolvedConfig } = this.resolveAliasChain( |
| context.model, |
| allAliases, |
| context.isChatModel, |
| ); |
|
|
| const modelToLevel = this.buildModelLevelMap(aliasChain, baseModel); |
| const allOverrides = [ |
| ...overrides, |
| ...customOverrides, |
| ...this.runtimeOverrides, |
| ]; |
| const matches = this.findMatchingOverrides( |
| allOverrides, |
| context, |
| modelToLevel, |
| ); |
|
|
| this.sortOverrides(matches); |
|
|
| let currentConfig: ModelConfig = { |
| model: baseModel, |
| generateContentConfig: resolvedConfig, |
| }; |
|
|
| for (const match of matches) { |
| currentConfig = ModelConfigService.merge( |
| currentConfig, |
| match.modelConfig, |
| ); |
| } |
|
|
| return { |
| model: currentConfig.model, |
| generateContentConfig: currentConfig.generateContentConfig ?? {}, |
| }; |
| } |
|
|
| private resolveAliasChain( |
| requestedModel: string, |
| allAliases: Record<string, ModelConfigAlias>, |
| isChatModel?: boolean, |
| ): { |
| aliasChain: string[]; |
| baseModel: string | undefined; |
| resolvedConfig: GenerateContentConfig; |
| } { |
| const aliasChain: string[] = []; |
|
|
| if (allAliases[requestedModel]) { |
| let current: string | undefined = requestedModel; |
| const visited = new Set<string>(); |
| while (current) { |
| const alias: ModelConfigAlias = allAliases[current]; |
| if (!alias) { |
| throw new Error(`Alias "${current}" not found.`); |
| } |
| if (visited.size >= MAX_ALIAS_CHAIN_DEPTH) { |
| throw new Error( |
| `Alias inheritance chain exceeded maximum depth of ${MAX_ALIAS_CHAIN_DEPTH}.`, |
| ); |
| } |
| if (visited.has(current)) { |
| throw new Error( |
| `Circular alias dependency: ${[...visited, current].join(' -> ')}`, |
| ); |
| } |
| visited.add(current); |
| aliasChain.push(current); |
| current = alias.extends; |
| } |
|
|
| |
| const reversedChain = [...aliasChain].reverse(); |
| let resolvedConfig: ModelConfig = {}; |
| for (const aliasName of reversedChain) { |
| const alias = allAliases[aliasName]; |
| resolvedConfig = ModelConfigService.merge( |
| resolvedConfig, |
| alias.modelConfig, |
| ); |
| } |
| return { |
| aliasChain: reversedChain, |
| baseModel: resolvedConfig.model, |
| resolvedConfig: resolvedConfig.generateContentConfig ?? {}, |
| }; |
| } |
|
|
| if (isChatModel) { |
| const fallbackAlias = 'chat-base'; |
| if (allAliases[fallbackAlias]) { |
| const fallbackResolution = this.resolveAliasChain( |
| fallbackAlias, |
| allAliases, |
| ); |
| return { |
| aliasChain: [...fallbackResolution.aliasChain, requestedModel], |
| baseModel: requestedModel, |
| resolvedConfig: fallbackResolution.resolvedConfig, |
| }; |
| } |
| } |
|
|
| return { |
| aliasChain: [requestedModel], |
| baseModel: requestedModel, |
| resolvedConfig: {}, |
| }; |
| } |
|
|
| private buildModelLevelMap( |
| aliasChain: string[], |
| baseModel: string | undefined, |
| ): Map<string, number> { |
| const modelToLevel = new Map<string, number>(); |
| |
| if (baseModel) { |
| modelToLevel.set(baseModel, 0); |
| } |
| |
| aliasChain.forEach((name, i) => modelToLevel.set(name, i + 1)); |
| return modelToLevel; |
| } |
|
|
| private findMatchingOverrides( |
| overrides: ModelConfigOverride[], |
| context: ModelConfigKey, |
| modelToLevel: Map<string, number>, |
| ): Array<{ |
| specificity: number; |
| level: number; |
| modelConfig: ModelConfig; |
| index: number; |
| }> { |
| return overrides |
| .map((override, index) => { |
| const matchEntries = Object.entries(override.match); |
| if (matchEntries.length === 0) return null; |
|
|
| let matchedLevel = 0; |
| const isMatch = matchEntries.every(([key, value]) => { |
| if (key === 'model') { |
| |
| const level = modelToLevel.get(value as string); |
| if (level === undefined) return false; |
| matchedLevel = level; |
| return true; |
| } |
| if (key === 'overrideScope' && value === 'core') { |
| return context.overrideScope === 'core' || !context.overrideScope; |
| } |
| |
| return context[key as keyof ModelConfigKey] === value; |
| }); |
|
|
| return isMatch |
| ? { |
| specificity: matchEntries.length, |
| level: matchedLevel, |
| modelConfig: override.modelConfig, |
| index, |
| } |
| : null; |
| }) |
| .filter((m): m is NonNullable<typeof m> => m !== null); |
| } |
|
|
| private sortOverrides( |
| matches: Array<{ specificity: number; level: number; index: number }>, |
| ): void { |
| matches.sort((a, b) => { |
| if (a.level !== b.level) { |
| return a.level - b.level; |
| } |
| if (a.specificity !== b.specificity) { |
| return a.specificity - b.specificity; |
| } |
| return a.index - b.index; |
| }); |
| } |
|
|
| getResolvedConfig(context: ModelConfigKey): ResolvedModelConfig { |
| const resolved = this.internalGetResolvedConfig(context); |
|
|
| if (!resolved.model) { |
| throw new Error( |
| `Could not resolve a model name for alias "${context.model}". Please ensure the alias chain or a matching override specifies a model.`, |
| ); |
| } |
|
|
| |
| return { |
| model: resolved.model, |
| generateContentConfig: resolved.generateContentConfig, |
| } as ResolvedModelConfig; |
| } |
|
|
| static isObject(item: unknown): item is Record<string, unknown> { |
| return !!item && typeof item === 'object' && !Array.isArray(item); |
| } |
|
|
| |
| |
| |
| |
| |
| static merge(base: ModelConfig, override: ModelConfig): ModelConfig { |
| return { |
| model: override.model ?? base.model, |
| generateContentConfig: ModelConfigService.deepMerge( |
| base.generateContentConfig, |
| override.generateContentConfig, |
| ), |
| }; |
| } |
|
|
| static deepMerge( |
| config1: GenerateContentConfig | undefined, |
| config2: GenerateContentConfig | undefined, |
| ): GenerateContentConfig { |
| return ModelConfigService.genericDeepMerge( |
| |
| config1 as Record<string, unknown> | undefined, |
| |
| config2 as Record<string, unknown> | undefined, |
| ) as GenerateContentConfig; |
| } |
|
|
| private static genericDeepMerge( |
| ...objects: Array<Record<string, unknown> | undefined> |
| ): Record<string, unknown> { |
| return objects.reduce((acc: Record<string, unknown>, obj) => { |
| if (!obj) { |
| return acc; |
| } |
|
|
| Object.keys(obj).forEach((key) => { |
| const accValue = acc[key]; |
| const objValue = obj[key]; |
|
|
| |
| |
| |
| |
| |
| if ( |
| ModelConfigService.isObject(accValue) && |
| ModelConfigService.isObject(objValue) |
| ) { |
| acc[key] = ModelConfigService.genericDeepMerge(accValue, objValue); |
| } else { |
| acc[key] = objValue; |
| } |
| }); |
|
|
| return acc; |
| }, {}); |
| } |
| } |
|
|