mirror of
https://github.com/supabase/supabase.git
synced 2026-10-05 09:25:06 +03:00
feat(assistant): per-endpoint reasoningEffort + model config cleanup (#43981)
We're exploring support for newer models like [gpt-5.4-nano](https://openai.com/index/introducing-gpt-5-4-mini-and-nano/) in Assistant. This model doesn't support the `'minimal'` reasoning effort level we use for gpt-5-mini which leads to vague errors. <img width="595" height="263" alt="CleanShot 2026-03-18 at 17 13 05@2x" src="https://github.com/user-attachments/assets/cf7c2370-322d-4a8a-be55-23e680db0aa0" /> Also, we've [previously discussed](https://supabase.slack.com/archives/C0161K73J1J/p1771544464850199?thread_ts=1771493920.775699&cid=C0161K73J1J) that reasoning adds unnecessary latency to otherwise simple AI completion endpoints like `title-v2`. We want more control of reasoning level independent of model/endpoint. This PR aims to solve both problems by: - making reasoning effort configurable on a per-request basis - adding compile-time guardrails to prevent selecting an incompatible reasoning level for models - adding a `DEFAULT_COMPLETION_MODEL` with minimal reasoning that we can update with newer models that support disabling reasoning (independent of Assistant chat model reasoning) Other improvements to our model config logic: - Fixes bug in `onboarding/design.ts` and `assistant.eval.ts` where `providerOptions` was being dropped - `getModel()` now returns a bundled `modelParams` object (spread into AI SDK calls) so `providerOptions` can't be accidentally omitted (this [has happened before](https://supabase.slack.com/archives/C0161K73J1J/p1771518443534309?thread_ts=1771493920.775699&cid=C0161K73J1J)) - Introduces an `ASSISTANT_MODELS` registry as a single source of truth for assistant model config, eliminating hardcoded model IDs across the codebase - Aligns free/pro model conditional logic with `assistant.advance_model` entitlement naming conventions instead of the `isLimited` pattern - Adds `console.error` logging of Assistant stream errors so we can interpret reasoning effort compatibility errors in the future (instead of just opaque "Sorry, I'm having trouble responding right now" card) - Removes unnecessary type casts and generally making the model config logic stricter - Removes pre-existing dead code: `anthropic` provider variant in `GetModelParams` / `PROVIDERS` registry that was never implemented in `getModel()` Now if you try to select an unsupported reasoning level you get a type error: <img width="1306" height="320" alt="CleanShot 2026-03-20 at 14 37 24@2x" src="https://github.com/user-attachments/assets/a6ac234b-5ea5-4d81-8e01-ac4be34a0800" /> And if for some reason an invalid reasoning level slips through, you now get a server-side error surfacing the issue: <img width="1268" height="204" alt="CleanShot 2026-03-20 at 14 58 14@2x" src="https://github.com/user-attachments/assets/aadc1b7a-9495-475f-9741-39979bd27cd7" /> I've tested gpt-5 and gpt-5-mini are still working on the staging preview and verified the models were selected properly in Braintrust logs. Both models are available on my Pro test account, and my Free test account shows the Pro upgrade CTA. Closes AI-446 Closes AI-551
This commit is contained in:
1 parent
e232c0b75e
commit
adf8b0c67c
22 files changed
+452
-281
No files matched your search
@@ -19,6 +19,12 @@ import { useOrgAiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi'
|
||||
import { useSelectedOrganizationQuery } from 'hooks/misc/useSelectedOrganization'
|
||||
import { useSelectedProjectQuery } from 'hooks/misc/useSelectedProject'
|
||||
import { useHotKey } from 'hooks/ui/useHotKey'
|
||||
import {
|
||||
DEFAULT_ASSISTANT_BASE_MODEL_ID,
|
||||
defaultAssistantModelId,
|
||||
isAssistantBaseModelId,
|
||||
isKnownAssistantModelId,
|
||||
} from 'lib/ai/model.utils'
|
||||
import { IS_PLATFORM } from 'lib/constants'
|
||||
import { uuidv4 } from 'lib/helpers'
|
||||
import type { AssistantModel } from 'state/ai-assistant-state'
|
||||
@@ -27,6 +33,7 @@ import { useSidebarManagerSnapshot } from 'state/sidebar-manager-state'
|
||||
import { useSqlEditorV2StateSnapshot } from 'state/sql-editor-v2'
|
||||
import { Button, cn, KeyboardShortcut } from 'ui'
|
||||
import { Admonition } from 'ui-patterns'
|
||||
|
||||
import { ButtonTooltip } from '../ButtonTooltip'
|
||||
import { ErrorBoundary } from '../ErrorBoundary/ErrorBoundary'
|
||||
import type { SqlSnippet } from './AIAssistant.types'
|
||||
@@ -72,14 +79,15 @@ export const AIAssistant = ({ className }: AIAssistantProps) => {
|
||||
const selectedModel = useMemo<AssistantModel>(() => {
|
||||
// While entitlements are loading, use the stored model without enforcing access
|
||||
if (isLoadingEntitlements) {
|
||||
return snap.model ?? 'gpt-5-mini'
|
||||
return snap.model ?? DEFAULT_ASSISTANT_BASE_MODEL_ID
|
||||
}
|
||||
|
||||
const defaultModel: AssistantModel = hasAccessToAdvanceModel ? 'gpt-5' : 'gpt-5-mini'
|
||||
const defaultModel = defaultAssistantModelId(hasAccessToAdvanceModel)
|
||||
const model = snap.model ?? defaultModel
|
||||
|
||||
if (!hasAccessToAdvanceModel && model === 'gpt-5') {
|
||||
return 'gpt-5-mini'
|
||||
if (!isKnownAssistantModelId(model)) return defaultModel
|
||||
if (!hasAccessToAdvanceModel && !isAssistantBaseModelId(model)) {
|
||||
return DEFAULT_ASSISTANT_BASE_MODEL_ID
|
||||
}
|
||||
|
||||
return model
|
||||
@@ -391,7 +399,7 @@ export const AIAssistant = ({ className }: AIAssistantProps) => {
|
||||
onCloseAssistant={() => closeSidebar(SIDEBAR_KEYS.AI_ASSISTANT)}
|
||||
showMetadataWarning={showMetadataWarning}
|
||||
updatedOptInSinceMCP={updatedOptInSinceMCP}
|
||||
isHipaaProjectDisallowed={isHipaaProjectDisallowed as boolean}
|
||||
isHipaaProjectDisallowed={isHipaaProjectDisallowed}
|
||||
aiOptInLevel={aiOptInLevel}
|
||||
/>
|
||||
{hasMessages ? (
|
||||
|
||||
@@ -5,6 +5,7 @@ import { useBreakpoint } from 'common'
|
||||
import { ExpandingTextArea } from 'ui'
|
||||
import { cn } from 'ui/src/lib/utils'
|
||||
import { ButtonTooltip } from '../ButtonTooltip'
|
||||
import type { AssistantModelId } from 'lib/ai/model.utils'
|
||||
import { type SqlSnippet } from './AIAssistant.types'
|
||||
import { ModelSelector } from './ModelSelector'
|
||||
import { getSnippetContent, SnippetRow } from './SnippetRow'
|
||||
@@ -45,9 +46,9 @@ export interface FormProps {
|
||||
/* If currently editing an existing message */
|
||||
isEditing?: boolean
|
||||
/* The currently selected AI model */
|
||||
selectedModel: 'gpt-5' | 'gpt-5-mini'
|
||||
selectedModel: AssistantModelId
|
||||
/* Callback when a model is chosen */
|
||||
onSelectModel: (model: 'gpt-5' | 'gpt-5-mini') => void
|
||||
onSelectModel: (model: AssistantModelId) => void
|
||||
}
|
||||
|
||||
const AssistantChatFormComponent = forwardRef<HTMLFormElement, FormProps>(
|
||||
|
||||
@@ -17,10 +17,12 @@ import {
|
||||
Tooltip,
|
||||
} from 'ui'
|
||||
import { useCheckEntitlements } from '@/hooks/misc/useCheckEntitlements'
|
||||
import { ASSISTANT_MODELS, isAdvanceOnlyModelId } from 'lib/ai/model.utils'
|
||||
import type { AssistantModelId } from 'lib/ai/model.utils'
|
||||
|
||||
interface ModelSelectorProps {
|
||||
selectedModel: 'gpt-5' | 'gpt-5-mini'
|
||||
onSelectModel: (model: 'gpt-5' | 'gpt-5-mini') => void
|
||||
selectedModel: AssistantModelId
|
||||
onSelectModel: (model: AssistantModelId) => void
|
||||
}
|
||||
|
||||
export const ModelSelector = ({ selectedModel, onSelectModel }: ModelSelectorProps) => {
|
||||
@@ -33,16 +35,19 @@ export const ModelSelector = ({ selectedModel, onSelectModel }: ModelSelectorPro
|
||||
|
||||
const slug = organization?.slug ?? '_'
|
||||
|
||||
const upgradeHref = `/org/${slug ?? '_'}/billing?panel=subscriptionPlan&source=ai-assistant-model`
|
||||
const upgradeHref = `/org/${slug}/billing?panel=subscriptionPlan&source=ai-assistant-model`
|
||||
|
||||
const handleSelectModel = (model: 'gpt-5' | 'gpt-5-mini') => {
|
||||
if (model === 'gpt-5' && !hasAccessToAdvanceModel) {
|
||||
const handleSelectModel = (modelId: AssistantModelId) => {
|
||||
if (isLoadingEntitlements && isAdvanceOnlyModelId(modelId)) {
|
||||
return
|
||||
}
|
||||
if (isAdvanceOnlyModelId(modelId) && !hasAccessToAdvanceModel) {
|
||||
setOpen(false)
|
||||
void router.push(upgradeHref)
|
||||
return
|
||||
}
|
||||
|
||||
onSelectModel(model)
|
||||
onSelectModel(modelId)
|
||||
setOpen(false)
|
||||
}
|
||||
|
||||
@@ -61,39 +66,35 @@ export const ModelSelector = ({ selectedModel, onSelectModel }: ModelSelectorPro
|
||||
<Command_Shadcn_>
|
||||
<CommandList_Shadcn_>
|
||||
<CommandGroup_Shadcn_>
|
||||
<CommandItem_Shadcn_
|
||||
value="gpt-5-mini"
|
||||
onSelect={() => handleSelectModel('gpt-5-mini')}
|
||||
className="flex justify-between"
|
||||
>
|
||||
<span>gpt-5-mini</span>
|
||||
{selectedModel === 'gpt-5-mini' && <Check className="h-3.5 w-3.5" />}
|
||||
</CommandItem_Shadcn_>
|
||||
<CommandItem_Shadcn_
|
||||
value="gpt-5"
|
||||
onSelect={() => handleSelectModel('gpt-5')}
|
||||
className="flex justify-between"
|
||||
>
|
||||
<span>gpt-5</span>
|
||||
{hasAccessToAdvanceModel ? (
|
||||
selectedModel === 'gpt-5' ? (
|
||||
<Check className="h-3.5 w-3.5" />
|
||||
) : null
|
||||
) : (
|
||||
<Tooltip>
|
||||
<TooltipTrigger asChild>
|
||||
<div>
|
||||
<Badge role="button" variant="warning">
|
||||
Upgrade
|
||||
</Badge>
|
||||
</div>
|
||||
</TooltipTrigger>
|
||||
<TooltipContent side="right">
|
||||
gpt-5 is available on Pro plans and above
|
||||
</TooltipContent>
|
||||
</Tooltip>
|
||||
)}
|
||||
</CommandItem_Shadcn_>
|
||||
{ASSISTANT_MODELS.map((m) => (
|
||||
<CommandItem_Shadcn_
|
||||
key={m.id}
|
||||
value={m.id}
|
||||
disabled={isLoadingEntitlements && isAdvanceOnlyModelId(m.id)}
|
||||
onSelect={() => handleSelectModel(m.id)}
|
||||
className="flex justify-between"
|
||||
>
|
||||
<span>{m.id}</span>
|
||||
{isAdvanceOnlyModelId(m.id) &&
|
||||
!hasAccessToAdvanceModel &&
|
||||
!isLoadingEntitlements ? (
|
||||
<Tooltip>
|
||||
<TooltipTrigger asChild>
|
||||
<div>
|
||||
<Badge role="button" variant="warning">
|
||||
Upgrade
|
||||
</Badge>
|
||||
</div>
|
||||
</TooltipTrigger>
|
||||
<TooltipContent side="right">
|
||||
{m.id} is available on Pro plans and above
|
||||
</TooltipContent>
|
||||
</Tooltip>
|
||||
) : (
|
||||
selectedModel === m.id && <Check className="h-3.5 w-3.5" />
|
||||
)}
|
||||
</CommandItem_Shadcn_>
|
||||
))}
|
||||
</CommandGroup_Shadcn_>
|
||||
</CommandList_Shadcn_>
|
||||
</Command_Shadcn_>
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import assert from 'node:assert'
|
||||
import { openai } from '@ai-sdk/openai'
|
||||
import { Eval } from 'braintrust'
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import { DEFAULT_ASSISTANT_BASE_MODEL_ID, getAssistantModelEntry } from 'lib/ai/model.utils'
|
||||
import { generateAssistantResponse } from 'lib/ai/generate-assistant-response'
|
||||
import { getMockTools } from 'lib/ai/tools/mock-tools'
|
||||
|
||||
@@ -25,9 +26,19 @@ Eval('Assistant', {
|
||||
trialCount: process.env.CI ? 3 : 1,
|
||||
data: () => dataset,
|
||||
task: async (input) => {
|
||||
const modelEntry = getAssistantModelEntry(DEFAULT_ASSISTANT_BASE_MODEL_ID)
|
||||
const modelResponse = await getModel({ provider: 'openai', modelEntry })
|
||||
if (modelResponse.error) throw modelResponse.error
|
||||
|
||||
const result = await generateAssistantResponse({
|
||||
model: openai('gpt-5-mini'),
|
||||
messages: [{ id: '1', role: 'user', parts: [{ type: 'text', text: input.prompt }] }],
|
||||
...modelResponse.modelParams,
|
||||
messages: [
|
||||
{
|
||||
id: '1',
|
||||
role: 'user',
|
||||
parts: [{ type: 'text', text: input.prompt }],
|
||||
},
|
||||
],
|
||||
tools: await getMockTools(input.mockTables ? { list_tables: input.mockTables } : undefined),
|
||||
})
|
||||
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import { openai } from '@ai-sdk/openai'
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import * as bedrockModule from './bedrock'
|
||||
import { getModel, ModelErrorMessage } from './model'
|
||||
import { getModel } from './model'
|
||||
import { DEFAULT_COMPLETION_MODEL, openaiModelEntry } from './model.utils'
|
||||
|
||||
vi.mock('@ai-sdk/openai', () => ({
|
||||
openai: vi.fn(() => 'openai-model'),
|
||||
@@ -9,7 +10,7 @@ vi.mock('@ai-sdk/openai', () => ({
|
||||
|
||||
vi.mock('./bedrock', async () => ({
|
||||
...(await vi.importActual('./bedrock')),
|
||||
createRoutedBedrock: vi.fn(() => async (modelId: string) => 'bedrock-model'),
|
||||
createRoutedBedrock: vi.fn(() => async (_modelId: string) => 'bedrock-model'),
|
||||
checkAwsCredentials: vi.fn(),
|
||||
}))
|
||||
|
||||
@@ -18,82 +19,81 @@ describe('getModel', () => {
|
||||
|
||||
beforeEach(() => {
|
||||
vi.resetAllMocks()
|
||||
vi.stubEnv('AWS_BEDROCK_ROLE_ARN', 'test')
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
process.env = { ...originalEnv }
|
||||
})
|
||||
|
||||
it('should return bedrock model and promptProviderOptions when default supports caching', async () => {
|
||||
it('returns bedrock model without promptProviderOptions', async () => {
|
||||
vi.mocked(bedrockModule.checkAwsCredentials).mockResolvedValue(true)
|
||||
vi.stubEnv('IS_THROTTLED', 'false')
|
||||
vi.stubEnv('AWS_BEDROCK_ROLE_ARN', 'test')
|
||||
|
||||
const { model, error, promptProviderOptions } = await getModel({
|
||||
const { modelParams, error, promptProviderOptions } = await getModel({
|
||||
provider: 'bedrock',
|
||||
routingKey: 'test',
|
||||
isLimited: false,
|
||||
})
|
||||
|
||||
expect(model).toEqual('bedrock-model')
|
||||
// Default bedrock model supportsCaching=false in registry, but if caller
|
||||
// specifies high-tier, provider options would be present
|
||||
expect(promptProviderOptions === undefined || typeof promptProviderOptions === 'object').toBe(
|
||||
true
|
||||
)
|
||||
expect(error).toBeUndefined()
|
||||
})
|
||||
|
||||
it('should return bedrock model when throttled (limited) with default model', async () => {
|
||||
vi.mocked(bedrockModule.checkAwsCredentials).mockResolvedValue(true)
|
||||
vi.stubEnv('IS_THROTTLED', 'true')
|
||||
|
||||
const { model, error, promptProviderOptions } = await getModel({
|
||||
routingKey: 'test',
|
||||
isLimited: true,
|
||||
})
|
||||
|
||||
expect(model).toEqual('bedrock-model')
|
||||
expect(modelParams?.model).toEqual('bedrock-model')
|
||||
expect(promptProviderOptions).toBeUndefined()
|
||||
expect(error).toBeUndefined()
|
||||
})
|
||||
|
||||
it('should return OpenAI model when AWS credentials are not available but OPENAI_API_KEY is set', async () => {
|
||||
it('returns error when bedrock credentials are not available', async () => {
|
||||
vi.mocked(bedrockModule.checkAwsCredentials).mockResolvedValue(false)
|
||||
process.env.OPENAI_API_KEY = 'test-key'
|
||||
|
||||
const { model, promptProviderOptions } = await getModel({
|
||||
routingKey: 'test',
|
||||
isLimited: false,
|
||||
const { error } = await getModel({ provider: 'bedrock', routingKey: 'test' })
|
||||
expect(error).toBeDefined()
|
||||
})
|
||||
|
||||
it('returns openai model with default model', async () => {
|
||||
vi.stubEnv('OPENAI_API_KEY', 'test-key')
|
||||
|
||||
const { modelParams, promptProviderOptions } = await getModel({
|
||||
provider: 'openai',
|
||||
modelEntry: openaiModelEntry({ id: 'gpt-5-mini' }),
|
||||
})
|
||||
|
||||
expect(model).toEqual('openai-model')
|
||||
// Default openai model in registry is gpt-5-mini
|
||||
expect(modelParams?.model).toEqual('openai-model')
|
||||
expect(openai).toHaveBeenCalledWith('gpt-5-mini')
|
||||
expect(promptProviderOptions).toBeUndefined()
|
||||
})
|
||||
|
||||
it('should return error when neither AWS credentials nor OPENAI_API_KEY is available', async () => {
|
||||
vi.mocked(bedrockModule.checkAwsCredentials).mockResolvedValue(false)
|
||||
delete process.env.OPENAI_API_KEY
|
||||
it('returns error when OPENAI_API_KEY is not available', async () => {
|
||||
vi.stubEnv('OPENAI_API_KEY', '')
|
||||
|
||||
const { error } = await getModel({ routingKey: 'test-key', isLimited: false })
|
||||
expect(error).toEqual(new Error(ModelErrorMessage))
|
||||
const { error } = await getModel({
|
||||
provider: 'openai',
|
||||
modelEntry: openaiModelEntry({ id: 'gpt-5-mini' }),
|
||||
})
|
||||
expect(error).toEqual(new Error('OPENAI_API_KEY not available'))
|
||||
})
|
||||
|
||||
it('returns specified provider and model when provided (openai gpt-5)', async () => {
|
||||
vi.mocked(bedrockModule.checkAwsCredentials).mockResolvedValue(false)
|
||||
process.env.OPENAI_API_KEY = 'test-key'
|
||||
process.env.IS_THROTTLED = 'false'
|
||||
it('returns openai gpt-5 when hasAccessToAdvanceModel and not throttled', async () => {
|
||||
vi.stubEnv('OPENAI_API_KEY', 'test-key')
|
||||
vi.stubEnv('IS_THROTTLED', 'false')
|
||||
|
||||
const { model, error } = await getModel({
|
||||
const { modelParams, error } = await getModel({
|
||||
provider: 'openai',
|
||||
model: 'gpt-5',
|
||||
routingKey: 'rk',
|
||||
isLimited: false,
|
||||
modelEntry: openaiModelEntry({ id: 'gpt-5', reasoningEffort: 'minimal' }),
|
||||
})
|
||||
|
||||
expect(error).toBeUndefined()
|
||||
expect(model).toEqual('openai-model')
|
||||
expect(modelParams?.model).toEqual('openai-model')
|
||||
expect(openai).toHaveBeenCalledWith('gpt-5')
|
||||
expect(modelParams?.providerOptions?.openai?.reasoningEffort).toBe('minimal')
|
||||
})
|
||||
|
||||
it('applies reasoningEffort from DEFAULT_COMPLETION_MODEL', async () => {
|
||||
vi.stubEnv('OPENAI_API_KEY', 'test-key')
|
||||
|
||||
const { modelParams, error } = await getModel({
|
||||
provider: 'openai',
|
||||
modelEntry: DEFAULT_COMPLETION_MODEL,
|
||||
})
|
||||
|
||||
expect(error).toBeUndefined()
|
||||
expect(openai).toHaveBeenCalledWith('gpt-5-mini')
|
||||
expect(modelParams?.providerOptions?.openai?.reasoningEffort).toBe('minimal')
|
||||
})
|
||||
})
|
||||
+54
-60
@@ -1,110 +1,104 @@
|
||||
import { openai } from '@ai-sdk/openai'
|
||||
import { LanguageModel } from 'ai'
|
||||
|
||||
import { checkAwsCredentials, createRoutedBedrock } from './bedrock'
|
||||
import {
|
||||
BedrockModel,
|
||||
Model,
|
||||
OpenAIModel,
|
||||
PROVIDERS,
|
||||
ProviderModelConfig,
|
||||
ProviderName,
|
||||
getDefaultModelForProvider,
|
||||
Model,
|
||||
OpenAIModelEntry,
|
||||
OpenAIModelId,
|
||||
ProviderModelConfig,
|
||||
PROVIDERS,
|
||||
} from './model.utils'
|
||||
|
||||
type PromptProviderOptions = Record<string, any>
|
||||
type ProviderOptions = Record<string, any>
|
||||
|
||||
type ModelSuccess = {
|
||||
model: LanguageModel
|
||||
/** Spread directly into AI SDK calls: `streamText({ ...modelParams, ... })` */
|
||||
modelParams: { model: LanguageModel; providerOptions?: ProviderOptions }
|
||||
promptProviderOptions?: PromptProviderOptions
|
||||
providerOptions?: ProviderOptions
|
||||
error?: never
|
||||
}
|
||||
|
||||
export type ModelError = {
|
||||
model?: never
|
||||
modelParams?: never
|
||||
promptProviderOptions?: never
|
||||
providerOptions?: never
|
||||
error: Error
|
||||
}
|
||||
|
||||
type ModelResponse = ModelSuccess | ModelError
|
||||
|
||||
export const ModelErrorMessage = 'No valid AI model available based on available credentials.'
|
||||
|
||||
export type GetModelParams = {
|
||||
provider?: ProviderName
|
||||
model?: Model
|
||||
routingKey: string
|
||||
isLimited?: boolean
|
||||
}
|
||||
export type GetModelParams =
|
||||
| {
|
||||
provider: 'openai'
|
||||
/**
|
||||
* Specifies which OpenAI model to use and its reasoning effort.
|
||||
* Create entries via `openaiModelEntry()` — reasoning effort is validated against the model
|
||||
* at compile time. Use `DEFAULT_COMPLETION_MODEL` for simple endpoints (minimal reasoning).
|
||||
* Callers are responsible for resolving the correct entry (including throttling/entitlement
|
||||
* fallbacks) before calling getModel.
|
||||
*/
|
||||
modelEntry: OpenAIModelEntry
|
||||
}
|
||||
| {
|
||||
provider: 'bedrock'
|
||||
/** Used for consistent hashing across Bedrock regions. */
|
||||
routingKey: string
|
||||
}
|
||||
|
||||
/**
|
||||
* Retrieves a LanguageModel from a specific provider and model.
|
||||
* - If provider/model not specified, auto-selects based on available credentials (prefers Bedrock).
|
||||
* - If isLimited is true, uses the provider's default model.
|
||||
* - Returns promptProviderOptions that callers can attach to the system message.
|
||||
* Retrieves a LanguageModel from a specific provider and model entry.
|
||||
* Callers are responsible for resolving the correct model entry (including throttling/entitlement
|
||||
* fallbacks) before calling this function.
|
||||
* Returns promptProviderOptions that callers can attach to the system message.
|
||||
*/
|
||||
export async function getModel({
|
||||
provider,
|
||||
model,
|
||||
routingKey,
|
||||
isLimited = true,
|
||||
}: GetModelParams): Promise<ModelResponse> {
|
||||
const envThrottled = process.env.IS_THROTTLED !== 'false'
|
||||
export async function getModel(params: GetModelParams): Promise<ModelResponse> {
|
||||
const { provider } = params
|
||||
|
||||
let preferredProvider: ProviderName | undefined = provider
|
||||
|
||||
const hasAwsCredentials = await checkAwsCredentials()
|
||||
const hasAwsBedrockRoleArn = !!process.env.AWS_BEDROCK_ROLE_ARN
|
||||
const hasOpenAIKey = !!process.env.OPENAI_API_KEY
|
||||
|
||||
// Auto-pick a provider if not specified defaulting to Bedrock
|
||||
if (!preferredProvider) {
|
||||
if (hasAwsBedrockRoleArn && hasAwsCredentials) {
|
||||
preferredProvider = 'bedrock'
|
||||
} else if (hasOpenAIKey) {
|
||||
preferredProvider = 'openai'
|
||||
}
|
||||
}
|
||||
|
||||
if (!preferredProvider) {
|
||||
return { error: new Error(ModelErrorMessage) }
|
||||
}
|
||||
|
||||
const providerRegistry = PROVIDERS[preferredProvider]
|
||||
const providerRegistry = PROVIDERS[provider]
|
||||
if (!providerRegistry) {
|
||||
return { error: new Error(`Unknown provider: ${preferredProvider}`) }
|
||||
return { error: new Error(`Unknown provider: ${provider}`) }
|
||||
}
|
||||
|
||||
const models = providerRegistry.models as Record<Model, ProviderModelConfig>
|
||||
const modelEntry = params.provider === 'openai' ? params.modelEntry : undefined
|
||||
|
||||
const useDefault = isLimited || envThrottled || !model || !models[model]
|
||||
const useDefault = !modelEntry?.id || !models[modelEntry.id]
|
||||
|
||||
const chosenModelId = useDefault ? getDefaultModelForProvider(preferredProvider) : model
|
||||
const chosenModelId = useDefault ? getDefaultModelForProvider(provider) : modelEntry?.id
|
||||
|
||||
if (preferredProvider === 'bedrock') {
|
||||
if (provider === 'bedrock') {
|
||||
const hasAwsCredentials = await checkAwsCredentials()
|
||||
const hasAwsBedrockRoleArn = !!process.env.AWS_BEDROCK_ROLE_ARN
|
||||
if (!hasAwsBedrockRoleArn || !hasAwsCredentials) {
|
||||
return { error: new Error('AWS Bedrock credentials not available') }
|
||||
}
|
||||
const bedrock = createRoutedBedrock(routingKey)
|
||||
const bedrock = createRoutedBedrock(params.routingKey)
|
||||
const model = await bedrock(chosenModelId as BedrockModel)
|
||||
const promptProviderOptions = (
|
||||
providerRegistry.models as Record<BedrockModel, ProviderModelConfig>
|
||||
)[chosenModelId as BedrockModel]?.promptProviderOptions
|
||||
return { model, promptProviderOptions }
|
||||
return { modelParams: { model }, promptProviderOptions }
|
||||
}
|
||||
|
||||
if (preferredProvider === 'openai') {
|
||||
if (!hasOpenAIKey) {
|
||||
if (provider === 'openai') {
|
||||
if (!process.env.OPENAI_API_KEY) {
|
||||
return { error: new Error('OPENAI_API_KEY not available') }
|
||||
}
|
||||
const baseProviderOptions = providerRegistry.providerOptions?.openai ?? {}
|
||||
const openaiProviderOptions = modelEntry?.reasoningEffort
|
||||
? { ...baseProviderOptions, reasoningEffort: modelEntry.reasoningEffort }
|
||||
: baseProviderOptions
|
||||
return {
|
||||
model: openai(chosenModelId as OpenAIModel),
|
||||
promptProviderOptions: models[chosenModelId as OpenAIModel]?.promptProviderOptions,
|
||||
providerOptions: providerRegistry.providerOptions,
|
||||
modelParams: {
|
||||
model: openai(chosenModelId as OpenAIModelId),
|
||||
providerOptions: { openai: openaiProviderOptions },
|
||||
},
|
||||
promptProviderOptions: models[chosenModelId as OpenAIModelId]?.promptProviderOptions,
|
||||
}
|
||||
}
|
||||
|
||||
return { error: new Error(`Unsupported provider: ${preferredProvider}`) }
|
||||
return { error: new Error(`Unsupported provider: ${provider}`) }
|
||||
}
|
||||
@@ -1,6 +1,19 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
import { getDefaultModelForProvider, PROVIDERS } from './model.utils'
|
||||
import {
|
||||
ASSISTANT_MODELS,
|
||||
DEFAULT_ASSISTANT_ADVANCE_MODEL_ID,
|
||||
DEFAULT_ASSISTANT_BASE_MODEL_ID,
|
||||
DEFAULT_COMPLETION_MODEL,
|
||||
defaultAssistantModelId,
|
||||
getAssistantModelEntry,
|
||||
getDefaultModelForProvider,
|
||||
isAdvanceOnlyModelId,
|
||||
isAssistantBaseModelId,
|
||||
isKnownAssistantModelId,
|
||||
openaiModelEntry,
|
||||
PROVIDERS,
|
||||
} from './model.utils'
|
||||
import type { ProviderName } from './model.utils'
|
||||
|
||||
describe('model.utils', () => {
|
||||
@@ -15,11 +28,6 @@ describe('model.utils', () => {
|
||||
expect(result).toBe('gpt-5-mini')
|
||||
})
|
||||
|
||||
it('should return correct default for anthropic provider', () => {
|
||||
const result = getDefaultModelForProvider('anthropic')
|
||||
expect(result).toBe('claude-3-5-haiku-20241022')
|
||||
})
|
||||
|
||||
it('should return undefined for unknown provider', () => {
|
||||
const result = getDefaultModelForProvider('unknown' as ProviderName)
|
||||
expect(result).toBeUndefined()
|
||||
@@ -43,15 +51,8 @@ describe('model.utils', () => {
|
||||
expect(Object.keys(PROVIDERS.openai.models)).toContain('gpt-5-mini')
|
||||
})
|
||||
|
||||
it('should have anthropic provider with models', () => {
|
||||
expect(PROVIDERS.anthropic).toBeDefined()
|
||||
expect(PROVIDERS.anthropic.models).toBeDefined()
|
||||
expect(Object.keys(PROVIDERS.anthropic.models)).toContain('claude-sonnet-4-20250514')
|
||||
expect(Object.keys(PROVIDERS.anthropic.models)).toContain('claude-3-5-haiku-20241022')
|
||||
})
|
||||
|
||||
it('should have exactly one default model per provider', () => {
|
||||
const providers: ProviderName[] = ['bedrock', 'openai', 'anthropic']
|
||||
const providers: ProviderName[] = ['bedrock', 'openai']
|
||||
|
||||
providers.forEach((provider) => {
|
||||
const models = PROVIDERS[provider].models
|
||||
@@ -61,11 +62,11 @@ describe('model.utils', () => {
|
||||
})
|
||||
|
||||
it('should have valid model configurations', () => {
|
||||
const providers: ProviderName[] = ['bedrock', 'openai', 'anthropic']
|
||||
const providers: ProviderName[] = ['bedrock', 'openai']
|
||||
|
||||
providers.forEach((provider) => {
|
||||
const models = PROVIDERS[provider].models
|
||||
Object.entries(models).forEach(([modelId, config]) => {
|
||||
Object.entries(models).forEach(([_modelId, config]) => {
|
||||
expect(config).toHaveProperty('default')
|
||||
expect(typeof config.default).toBe('boolean')
|
||||
})
|
||||
@@ -76,13 +77,83 @@ describe('model.utils', () => {
|
||||
const sonnetModel = PROVIDERS.bedrock.models['anthropic.claude-3-7-sonnet-20250219-v1:0']
|
||||
expect(sonnetModel.promptProviderOptions).toBeDefined()
|
||||
expect(sonnetModel.promptProviderOptions?.bedrock).toBeDefined()
|
||||
expect(sonnetModel.promptProviderOptions?.bedrock?.cachePoint).toEqual({ type: 'default' })
|
||||
expect(sonnetModel.promptProviderOptions?.bedrock?.cachePoint).toEqual({
|
||||
type: 'default',
|
||||
})
|
||||
})
|
||||
|
||||
it('should have openai provider with providerOptions', () => {
|
||||
expect(PROVIDERS.openai.providerOptions).toBeDefined()
|
||||
expect(PROVIDERS.openai.providerOptions?.openai).toBeDefined()
|
||||
expect(PROVIDERS.openai.providerOptions?.openai?.reasoningEffort).toBe('minimal')
|
||||
expect(PROVIDERS.openai.providerOptions?.openai?.reasoningEffort).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe('assistant model registry', () => {
|
||||
it('should have non-empty base and advance tiers', () => {
|
||||
expect(
|
||||
ASSISTANT_MODELS.filter((m) => !m.requiresAdvanceModelEntitlement).length
|
||||
).toBeGreaterThan(0)
|
||||
expect(
|
||||
ASSISTANT_MODELS.filter((m) => m.requiresAdvanceModelEntitlement).length
|
||||
).toBeGreaterThan(0)
|
||||
})
|
||||
|
||||
it('all model IDs should be unique', () => {
|
||||
const ids = ASSISTANT_MODELS.map((m) => m.id)
|
||||
expect(new Set(ids).size).toBe(ids.length)
|
||||
})
|
||||
|
||||
it('should have all models in openai provider registry', () => {
|
||||
ASSISTANT_MODELS.forEach((entry) => {
|
||||
expect(Object.keys(PROVIDERS.openai.models)).toContain(entry.id)
|
||||
})
|
||||
})
|
||||
|
||||
it('defaults should satisfy unions', () => {
|
||||
expect(DEFAULT_ASSISTANT_BASE_MODEL_ID).toBe('gpt-5-mini')
|
||||
expect(DEFAULT_ASSISTANT_ADVANCE_MODEL_ID).toBe('gpt-5')
|
||||
expect(defaultAssistantModelId(false)).toBe(DEFAULT_ASSISTANT_BASE_MODEL_ID)
|
||||
expect(defaultAssistantModelId(true)).toBe(DEFAULT_ASSISTANT_ADVANCE_MODEL_ID)
|
||||
})
|
||||
|
||||
it('isAssistantBaseModelId / isAdvanceOnlyModelId', () => {
|
||||
expect(isAssistantBaseModelId('gpt-5-mini')).toBe(true)
|
||||
expect(isAssistantBaseModelId('gpt-5')).toBe(false)
|
||||
expect(isAdvanceOnlyModelId('gpt-5')).toBe(true)
|
||||
expect(isAdvanceOnlyModelId('gpt-5-mini')).toBe(false)
|
||||
})
|
||||
|
||||
it('isKnownAssistantModelId', () => {
|
||||
expect(isKnownAssistantModelId('gpt-5-mini')).toBe(true)
|
||||
expect(isKnownAssistantModelId('gpt-5')).toBe(true)
|
||||
expect(isKnownAssistantModelId('unknown')).toBe(false)
|
||||
})
|
||||
|
||||
it('getAssistantModelEntry returns config for known ids', () => {
|
||||
expect(getAssistantModelEntry('gpt-5-mini').reasoningEffort).toBe('minimal')
|
||||
expect(getAssistantModelEntry('gpt-5').reasoningEffort).toBe('minimal')
|
||||
expect(getAssistantModelEntry('gpt-5-mini')).toEqual(
|
||||
ASSISTANT_MODELS.find((m) => m.id === 'gpt-5-mini')
|
||||
)
|
||||
})
|
||||
|
||||
it('DEFAULT_COMPLETION_MODEL is gpt-5-mini with minimal reasoning effort', () => {
|
||||
expect(DEFAULT_COMPLETION_MODEL.id).toBe(DEFAULT_ASSISTANT_BASE_MODEL_ID)
|
||||
expect(DEFAULT_COMPLETION_MODEL.reasoningEffort).toBe('minimal')
|
||||
})
|
||||
|
||||
it('openaiModelEntry enforces valid reasoning effort at compile time', () => {
|
||||
// Valid: supported effort level
|
||||
const withEffort = openaiModelEntry({
|
||||
id: 'gpt-5-mini',
|
||||
reasoningEffort: 'low',
|
||||
})
|
||||
expect(withEffort.reasoningEffort).toBe('low')
|
||||
|
||||
// Valid: no effort
|
||||
const withoutEffort = openaiModelEntry({ id: 'gpt-5-mini' })
|
||||
expect(withoutEffort.reasoningEffort).toBeUndefined()
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -1,12 +1,115 @@
|
||||
export type ProviderName = 'bedrock' | 'openai' | 'anthropic'
|
||||
export type ProviderName = 'bedrock' | 'openai'
|
||||
|
||||
export type BedrockModel = 'anthropic.claude-3-7-sonnet-20250219-v1:0' | 'openai.gpt-oss-120b-1:0'
|
||||
|
||||
export type OpenAIModel = 'gpt-5' | 'gpt-5-mini'
|
||||
export type OpenAIModelId = 'gpt-5' | 'gpt-5-mini'
|
||||
|
||||
export type AnthropicModel = 'claude-sonnet-4-20250514' | 'claude-3-5-haiku-20241022'
|
||||
// Source: https://developers.openai.com/api/docs/guides/reasoning + per-model pages
|
||||
export type ReasoningEffort = 'none' | 'minimal' | 'low' | 'medium' | 'high' | 'xhigh'
|
||||
|
||||
export type Model = BedrockModel | OpenAIModel | AnthropicModel
|
||||
// Per-model reasoning effort compatibility.
|
||||
// When adding a model, verify supported levels in the community matrix and add an entry:
|
||||
// https://community.openai.com/t/request-for-compatibility-matrix-reasoning-effort-sampling-parameters-across-gpt-5-series/1371738/2
|
||||
type ModelReasoningSupport = {
|
||||
'gpt-5': 'minimal' | 'low' | 'medium' | 'high'
|
||||
'gpt-5-mini': 'minimal' | 'low' | 'medium' | 'high'
|
||||
}
|
||||
|
||||
type ReasoningEffortFor<ModelId extends OpenAIModelId> = ModelId extends keyof ModelReasoningSupport
|
||||
? ModelReasoningSupport[ModelId]
|
||||
: never
|
||||
|
||||
/** Type-safe factory for configuring OpenAI models with compatible reasoning efforts. */
|
||||
export function openaiModelEntry<
|
||||
ModelId extends OpenAIModelId,
|
||||
RequiresAdvance extends boolean = false,
|
||||
>(config: {
|
||||
id: ModelId
|
||||
/** When true, the model requires the `assistant.advance_model` entitlement (paid plans). Defaults to false. */
|
||||
requiresAdvanceModelEntitlement?: RequiresAdvance
|
||||
/**
|
||||
* When omitted, OpenAI applies its own default reasoning effort for the model,
|
||||
* which may not be zero. Use an explicit level to control cost and latency.
|
||||
*/
|
||||
reasoningEffort?: ReasoningEffortFor<ModelId>
|
||||
}): {
|
||||
id: ModelId
|
||||
requiresAdvanceModelEntitlement: RequiresAdvance
|
||||
reasoningEffort?: ReasoningEffortFor<ModelId>
|
||||
} {
|
||||
return {
|
||||
requiresAdvanceModelEntitlement: false as RequiresAdvance,
|
||||
...config,
|
||||
}
|
||||
}
|
||||
|
||||
export type OpenAIModelEntry = ReturnType<typeof openaiModelEntry>
|
||||
|
||||
/** Default model entry for simple completion endpoints where latency is more important than reasoning. */
|
||||
export const DEFAULT_COMPLETION_MODEL = openaiModelEntry({
|
||||
id: 'gpt-5-mini',
|
||||
reasoningEffort: 'minimal',
|
||||
})
|
||||
|
||||
// Single source of truth for all Assistant chat model variants and their reasoning levels.
|
||||
// Models with requiresAdvanceModelEntitlement false are available to all users; true requires the assistant.advance_model entitlement.
|
||||
export const ASSISTANT_MODELS = [
|
||||
openaiModelEntry({
|
||||
id: 'gpt-5-mini',
|
||||
requiresAdvanceModelEntitlement: false,
|
||||
reasoningEffort: 'minimal',
|
||||
}),
|
||||
openaiModelEntry({
|
||||
id: 'gpt-5',
|
||||
requiresAdvanceModelEntitlement: true,
|
||||
reasoningEffort: 'minimal',
|
||||
}),
|
||||
] as const
|
||||
|
||||
export type AssistantBaseModelId = Extract<
|
||||
(typeof ASSISTANT_MODELS)[number],
|
||||
{ requiresAdvanceModelEntitlement: false }
|
||||
>['id']
|
||||
export type AssistantModelId = (typeof ASSISTANT_MODELS)[number]['id']
|
||||
|
||||
const ASSISTANT_MODELS_MAP = Object.fromEntries(ASSISTANT_MODELS.map((m) => [m.id, m])) as Record<
|
||||
AssistantModelId,
|
||||
(typeof ASSISTANT_MODELS)[number]
|
||||
>
|
||||
|
||||
export const DEFAULT_ASSISTANT_BASE_MODEL_ID = 'gpt-5-mini' satisfies AssistantBaseModelId
|
||||
|
||||
export const DEFAULT_ASSISTANT_ADVANCE_MODEL_ID = 'gpt-5' satisfies AssistantModelId
|
||||
|
||||
export function defaultAssistantModelId(hasAccessToAdvanceModel: boolean): AssistantModelId {
|
||||
return hasAccessToAdvanceModel
|
||||
? DEFAULT_ASSISTANT_ADVANCE_MODEL_ID
|
||||
: DEFAULT_ASSISTANT_BASE_MODEL_ID
|
||||
}
|
||||
|
||||
export function isKnownAssistantModelId(id: string): id is AssistantModelId {
|
||||
return Object.hasOwn(ASSISTANT_MODELS_MAP, id)
|
||||
}
|
||||
|
||||
export function isAssistantBaseModelId(id: string): id is AssistantBaseModelId {
|
||||
return (
|
||||
id in ASSISTANT_MODELS_MAP &&
|
||||
!ASSISTANT_MODELS_MAP[id as AssistantModelId].requiresAdvanceModelEntitlement
|
||||
)
|
||||
}
|
||||
|
||||
export function isAdvanceOnlyModelId(id: string): boolean {
|
||||
return (
|
||||
id in ASSISTANT_MODELS_MAP &&
|
||||
ASSISTANT_MODELS_MAP[id as AssistantModelId].requiresAdvanceModelEntitlement
|
||||
)
|
||||
}
|
||||
|
||||
export function getAssistantModelEntry(id: AssistantModelId): (typeof ASSISTANT_MODELS)[number] {
|
||||
return ASSISTANT_MODELS_MAP[id]
|
||||
}
|
||||
|
||||
export type Model = BedrockModel | OpenAIModelId
|
||||
|
||||
export type ProviderModelConfig = {
|
||||
/** Optional providerOptions to attach to the system message for this model */
|
||||
@@ -21,11 +124,7 @@ export type ProviderRegistry = {
|
||||
providerOptions?: Record<string, any>
|
||||
}
|
||||
openai: {
|
||||
models: Record<OpenAIModel, ProviderModelConfig>
|
||||
providerOptions?: Record<string, any>
|
||||
}
|
||||
anthropic: {
|
||||
models: Record<AnthropicModel, ProviderModelConfig>
|
||||
models: Record<OpenAIModelId, ProviderModelConfig>
|
||||
providerOptions?: Record<string, any>
|
||||
}
|
||||
}
|
||||
@@ -54,17 +153,10 @@ export const PROVIDERS: ProviderRegistry = {
|
||||
},
|
||||
providerOptions: {
|
||||
openai: {
|
||||
reasoningEffort: 'minimal',
|
||||
store: false,
|
||||
},
|
||||
},
|
||||
},
|
||||
anthropic: {
|
||||
models: {
|
||||
'claude-sonnet-4-20250514': { default: false },
|
||||
'claude-3-5-haiku-20241022': { default: true },
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
export function getDefaultModelForProvider(provider: ProviderName): Model | undefined {
|
||||
|
||||
@@ -102,7 +102,7 @@ describe('ai/org-ai-details', () => {
|
||||
})
|
||||
})
|
||||
|
||||
it('should return AI opt-in level and limited status', async () => {
|
||||
it('should return AI opt-in level and assistant advance-model flag', async () => {
|
||||
const mockOrg = {
|
||||
id: 1,
|
||||
slug: 'test-org',
|
||||
@@ -126,14 +126,14 @@ describe('ai/org-ai-details', () => {
|
||||
|
||||
expect(result).toEqual({
|
||||
aiOptInLevel: 'schema_only',
|
||||
isLimited: true,
|
||||
hasAccessToAdvanceModel: false,
|
||||
isHipaaEnabled: false,
|
||||
orgId: 1,
|
||||
planId: 'free',
|
||||
})
|
||||
})
|
||||
|
||||
it('should mark pro plan as not limited', async () => {
|
||||
it('should set hasAccessToAdvanceModel when entitlement grants access', async () => {
|
||||
const mockOrg = {
|
||||
id: 1,
|
||||
slug: 'test-org',
|
||||
@@ -155,7 +155,7 @@ describe('ai/org-ai-details', () => {
|
||||
projectRef: 'test-project',
|
||||
})
|
||||
|
||||
expect(result.isLimited).toBe(false)
|
||||
expect(result.hasAccessToAdvanceModel).toBe(true)
|
||||
})
|
||||
|
||||
it('should throw error when project and org do not match', async () => {
|
||||
@@ -277,7 +277,7 @@ describe('ai/org-ai-details', () => {
|
||||
projectRef: 'test-project',
|
||||
})
|
||||
|
||||
expect(result.isLimited).toBe(false) // Has advance model entitlement
|
||||
expect(result.hasAccessToAdvanceModel).toBe(true)
|
||||
})
|
||||
|
||||
it('should return isHipaaEnabled true when subscription has HIPAA addon and project is sensitive', async () => {
|
||||
|
||||
@@ -37,12 +37,12 @@ export const getOrgAIDetails = async ({
|
||||
}
|
||||
|
||||
const aiOptInLevel = getAiOptInLevel(selectedOrg?.opt_in_tags)
|
||||
const isLimited = !advanceModelAccess.hasAccess
|
||||
const hasAccessToAdvanceModel = advanceModelAccess.hasAccess
|
||||
const isHipaaEnabled = subscriptionHasHipaaAddon(subscription) && !!projectSettings?.is_sensitive
|
||||
|
||||
return {
|
||||
aiOptInLevel,
|
||||
isLimited,
|
||||
hasAccessToAdvanceModel,
|
||||
isHipaaEnabled,
|
||||
orgId: selectedOrg?.id,
|
||||
planId: selectedOrg?.plan.id,
|
||||
|
||||
@@ -47,16 +47,14 @@ test('generateV4 calls the tool sanitizer', async () => {
|
||||
vi.mock('lib/ai/org-ai-details', () => ({
|
||||
getOrgAIDetails: vi.fn().mockResolvedValue({
|
||||
aiOptInLevel: 'schema_and_log_and_data',
|
||||
isLimited: false,
|
||||
hasAccessToAdvanceModel: true,
|
||||
}),
|
||||
}))
|
||||
|
||||
vi.mock('lib/ai/model', () => ({
|
||||
getModel: vi.fn().mockResolvedValue({
|
||||
model: {},
|
||||
error: null,
|
||||
modelParams: { model: {} },
|
||||
promptProviderOptions: {},
|
||||
providerOptions: {},
|
||||
}),
|
||||
}))
|
||||
|
||||
|
||||
@@ -45,14 +45,13 @@ test('rate calls the tool sanitizer', async () => {
|
||||
vi.mock('lib/ai/org-ai-details', () => ({
|
||||
getOrgAIDetails: vi.fn().mockResolvedValue({
|
||||
aiOptInLevel: 'schema_and_log_and_data',
|
||||
isLimited: false,
|
||||
hasAccessToAdvanceModel: true,
|
||||
}),
|
||||
}))
|
||||
|
||||
vi.mock('lib/ai/model', () => ({
|
||||
getModel: vi.fn().mockResolvedValue({
|
||||
model: {},
|
||||
error: null,
|
||||
modelParams: { model: {} },
|
||||
}),
|
||||
}))
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ import { source } from 'common-tags'
|
||||
import { executeSql } from 'data/sql/execute-sql-query'
|
||||
import { AiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi'
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils'
|
||||
import { getOrgAIDetails } from 'lib/ai/org-ai-details'
|
||||
import {
|
||||
EDGE_FUNCTION_PROMPT,
|
||||
@@ -65,13 +66,12 @@ async function handler(req: NextApiRequest, res: NextApiResponse) {
|
||||
|
||||
// For code completion, we always use the limited model
|
||||
const {
|
||||
model,
|
||||
modelParams,
|
||||
error: modelError,
|
||||
promptProviderOptions,
|
||||
providerOptions,
|
||||
} = await getModel({
|
||||
provider: 'openai',
|
||||
routingKey: projectRef,
|
||||
modelEntry: DEFAULT_COMPLETION_MODEL,
|
||||
})
|
||||
|
||||
if (modelError) {
|
||||
@@ -155,8 +155,7 @@ async function handler(req: NextApiRequest, res: NextApiResponse) {
|
||||
})
|
||||
|
||||
const { text } = await generateText({
|
||||
model,
|
||||
providerOptions,
|
||||
...modelParams,
|
||||
stopWhen: stepCountIs(5),
|
||||
messages: coreMessages,
|
||||
tools,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { generateText, Output } from 'ai'
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils'
|
||||
import apiWrapper from 'lib/api/apiWrapper'
|
||||
import { NextApiRequest, NextApiResponse } from 'next'
|
||||
import { z } from 'zod'
|
||||
@@ -28,13 +29,9 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
}
|
||||
|
||||
try {
|
||||
const {
|
||||
model,
|
||||
error: modelError,
|
||||
providerOptions,
|
||||
} = await getModel({
|
||||
const { modelParams, error: modelError } = await getModel({
|
||||
provider: 'openai',
|
||||
routingKey: 'feedback',
|
||||
modelEntry: DEFAULT_COMPLETION_MODEL,
|
||||
})
|
||||
|
||||
if (modelError) {
|
||||
@@ -42,8 +39,7 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
}
|
||||
|
||||
const { output } = await generateText({
|
||||
model,
|
||||
providerOptions,
|
||||
...modelParams,
|
||||
output: Output.object({
|
||||
schema: z.object({
|
||||
feedback_category: z.enum(['support', 'feedback', 'unknown']),
|
||||
|
||||
@@ -5,6 +5,7 @@ import { rateMessageResponseSchema } from 'components/ui/AIAssistantPanel/Messag
|
||||
import type { AiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi'
|
||||
import { IS_TRACING_ENABLED } from 'lib/ai/braintrust-logger'
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils'
|
||||
import { getOrgAIDetails } from 'lib/ai/org-ai-details'
|
||||
import { sanitizeMessagePart } from 'lib/ai/tools/tool-sanitizer'
|
||||
import apiWrapper from 'lib/api/apiWrapper'
|
||||
@@ -96,14 +97,9 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
})
|
||||
|
||||
try {
|
||||
const {
|
||||
model,
|
||||
error: modelError,
|
||||
providerOptions,
|
||||
} = await getModel({
|
||||
const { modelParams, error: modelError } = await getModel({
|
||||
provider: 'openai',
|
||||
isLimited: true,
|
||||
routingKey: 'feedback',
|
||||
modelEntry: DEFAULT_COMPLETION_MODEL,
|
||||
})
|
||||
|
||||
if (modelError) {
|
||||
@@ -111,8 +107,7 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
}
|
||||
|
||||
const { output } = await generateText({
|
||||
model,
|
||||
providerOptions,
|
||||
...modelParams,
|
||||
output: Output.object({ schema: rateMessageResponseSchema }),
|
||||
prompt: `
|
||||
Your job is to look at a Supabase Assistant conversation, which the user has given feedback on, and classify it.
|
||||
|
||||
@@ -4,6 +4,7 @@ import { NextApiRequest, NextApiResponse } from 'next'
|
||||
import { z } from 'zod'
|
||||
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils'
|
||||
import apiWrapper from 'lib/api/apiWrapper'
|
||||
|
||||
export const maxDuration = 60
|
||||
@@ -64,9 +65,9 @@ const wrapper = (req: NextApiRequest, res: NextApiResponse) =>
|
||||
export default wrapper
|
||||
|
||||
async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
const { model, error: modelError } = await getModel({
|
||||
const { modelParams, error: modelError } = await getModel({
|
||||
provider: 'openai',
|
||||
routingKey: 'onboarding',
|
||||
modelEntry: DEFAULT_COMPLETION_MODEL,
|
||||
})
|
||||
|
||||
if (modelError) {
|
||||
@@ -76,7 +77,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
const { messages } = req.body
|
||||
|
||||
const result = streamText({
|
||||
model,
|
||||
...modelParams,
|
||||
system: source`
|
||||
You are a Supabase expert who helps people set up their Supabase project. You specializes in database schema design. You are to help the user design a database schema for their application but also suggest Supabase services they should use.
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { generateText, Output } from 'ai'
|
||||
import { source } from 'common-tags'
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils'
|
||||
import apiWrapper from 'lib/api/apiWrapper'
|
||||
import { NextApiRequest, NextApiResponse } from 'next'
|
||||
import { z } from 'zod'
|
||||
@@ -33,13 +34,9 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
}
|
||||
|
||||
try {
|
||||
const {
|
||||
model,
|
||||
error: modelError,
|
||||
providerOptions,
|
||||
} = await getModel({
|
||||
const { modelParams, error: modelError } = await getModel({
|
||||
provider: 'openai',
|
||||
routingKey: 'cron',
|
||||
modelEntry: DEFAULT_COMPLETION_MODEL,
|
||||
})
|
||||
|
||||
if (modelError) {
|
||||
@@ -47,8 +44,7 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
}
|
||||
|
||||
const result = await generateText({
|
||||
model,
|
||||
providerOptions,
|
||||
...modelParams,
|
||||
output: Output.object({ schema: cronSchema }),
|
||||
prompt: source`
|
||||
You are a cron syntax expert. Your purpose is to convert natural language time descriptions into valid cron expressions for pg_cron.
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { generateText, Output } from 'ai'
|
||||
import { source } from 'common-tags'
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils'
|
||||
import apiWrapper from 'lib/api/apiWrapper'
|
||||
import {
|
||||
filterGroupSchemaForAI,
|
||||
@@ -34,13 +35,9 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
const { prompt, filterProperties } = parseResult.data
|
||||
|
||||
try {
|
||||
const {
|
||||
model,
|
||||
error: modelError,
|
||||
providerOptions,
|
||||
} = await getModel({
|
||||
const { modelParams, error: modelError } = await getModel({
|
||||
provider: 'openai',
|
||||
routingKey: 'sql',
|
||||
modelEntry: DEFAULT_COMPLETION_MODEL,
|
||||
})
|
||||
|
||||
if (modelError) {
|
||||
@@ -62,8 +59,7 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
}))
|
||||
|
||||
const result = await generateText({
|
||||
model,
|
||||
providerOptions,
|
||||
...modelParams,
|
||||
output: Output.object({ schema: filterGroupSchemaForAI }),
|
||||
prompt: source`
|
||||
You are an expert Postgres filter builder. Convert the user's request into structured filters.
|
||||
|
||||
@@ -6,6 +6,14 @@ import { executeSql } from 'data/sql/execute-sql-query'
|
||||
import type { AiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi'
|
||||
import { generateAssistantResponse } from 'lib/ai/generate-assistant-response'
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import {
|
||||
DEFAULT_ASSISTANT_ADVANCE_MODEL_ID,
|
||||
DEFAULT_ASSISTANT_BASE_MODEL_ID,
|
||||
getAssistantModelEntry,
|
||||
isAssistantBaseModelId,
|
||||
isKnownAssistantModelId,
|
||||
type AssistantModelId,
|
||||
} from 'lib/ai/model.utils'
|
||||
import { getOrgAIDetails } from 'lib/ai/org-ai-details'
|
||||
import { getTools } from 'lib/ai/tools'
|
||||
import apiWrapper from 'lib/api/apiWrapper'
|
||||
@@ -53,7 +61,7 @@ const requestBodySchema = z.object({
|
||||
chatId: z.string().optional(),
|
||||
chatName: z.string().optional(),
|
||||
orgSlug: z.string().optional(),
|
||||
model: z.enum(['gpt-5', 'gpt-5-mini']).optional(),
|
||||
model: z.string().optional(),
|
||||
})
|
||||
|
||||
async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: JwtPayload) {
|
||||
@@ -80,25 +88,32 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw
|
||||
orgSlug,
|
||||
chatId,
|
||||
chatName,
|
||||
model: requestedModel,
|
||||
model: rawRequestedModel,
|
||||
} = data
|
||||
|
||||
const messagesValidation = await safeValidateUIMessages({ messages: rawMessages })
|
||||
const requestedModel: AssistantModelId | undefined =
|
||||
rawRequestedModel && isKnownAssistantModelId(rawRequestedModel) ? rawRequestedModel : undefined
|
||||
|
||||
const messagesValidation = await safeValidateUIMessages({
|
||||
messages: rawMessages,
|
||||
})
|
||||
if (!messagesValidation.success) {
|
||||
return res
|
||||
.status(400)
|
||||
.json({ error: 'Invalid request body', message: messagesValidation.error.message })
|
||||
return res.status(400).json({
|
||||
error: 'Invalid request body',
|
||||
message: messagesValidation.error.message,
|
||||
})
|
||||
}
|
||||
const messages = messagesValidation.data
|
||||
|
||||
let aiOptInLevel: AiOptInLevel = 'disabled'
|
||||
let isLimited = false
|
||||
let hasAccessToAdvanceModel = false
|
||||
let isHipaaEnabled = false
|
||||
let orgId: number | undefined
|
||||
let planId: string | undefined
|
||||
|
||||
if (!IS_PLATFORM) {
|
||||
aiOptInLevel = 'schema'
|
||||
hasAccessToAdvanceModel = true
|
||||
}
|
||||
|
||||
if (IS_PLATFORM && orgSlug && authorization && projectRef) {
|
||||
@@ -106,7 +121,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw
|
||||
// Get organizations and compute opt in level server-side
|
||||
const {
|
||||
aiOptInLevel: orgAIOptInLevel,
|
||||
isLimited: orgAILimited,
|
||||
hasAccessToAdvanceModel: orgHasAccessToAdvanceModel,
|
||||
isHipaaEnabled: orgIsHipaaEnabled,
|
||||
orgId: fetchedOrgId,
|
||||
planId: fetchedPlanId,
|
||||
@@ -117,7 +132,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw
|
||||
})
|
||||
|
||||
aiOptInLevel = orgAIOptInLevel
|
||||
isLimited = orgAILimited
|
||||
hasAccessToAdvanceModel = orgHasAccessToAdvanceModel
|
||||
isHipaaEnabled = orgIsHipaaEnabled
|
||||
orgId = fetchedOrgId
|
||||
planId = fetchedPlanId
|
||||
@@ -128,16 +143,20 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw
|
||||
}
|
||||
}
|
||||
|
||||
const envThrottled = process.env.IS_THROTTLED !== 'false'
|
||||
|
||||
let effectiveModel: AssistantModelId = requestedModel ?? DEFAULT_ASSISTANT_ADVANCE_MODEL_ID
|
||||
if (!hasAccessToAdvanceModel || (envThrottled && !isAssistantBaseModelId(effectiveModel))) {
|
||||
effectiveModel = DEFAULT_ASSISTANT_BASE_MODEL_ID
|
||||
}
|
||||
|
||||
const {
|
||||
model,
|
||||
modelParams,
|
||||
error: modelError,
|
||||
promptProviderOptions,
|
||||
providerOptions,
|
||||
} = await getModel({
|
||||
provider: 'openai',
|
||||
model: requestedModel ?? 'gpt-5',
|
||||
routingKey: projectRef,
|
||||
isLimited,
|
||||
modelEntry: getAssistantModelEntry(effectiveModel),
|
||||
})
|
||||
|
||||
if (modelError) {
|
||||
@@ -184,7 +203,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw
|
||||
|
||||
const result = await generateAssistantResponse({
|
||||
messages,
|
||||
model,
|
||||
...modelParams,
|
||||
tools,
|
||||
aiOptInLevel,
|
||||
getSchemas: aiOptInLevel !== 'disabled' ? getSchemas : undefined,
|
||||
@@ -197,7 +216,6 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw
|
||||
planId,
|
||||
requestedModel,
|
||||
promptProviderOptions,
|
||||
providerOptions,
|
||||
abortSignal: abortController.signal,
|
||||
onSpanCreated: (spanId) => {
|
||||
res.setHeader('x-braintrust-span-id', spanId)
|
||||
@@ -208,6 +226,8 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw
|
||||
sendReasoning: true,
|
||||
headers: { 'Content-Encoding': 'none' },
|
||||
onError: (error) => {
|
||||
console.error('Assistant stream error:', error)
|
||||
|
||||
if (error == null) {
|
||||
return 'unknown error'
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ import { IS_PLATFORM } from 'common'
|
||||
import { source } from 'common-tags'
|
||||
import type { AiOptInLevel } from 'hooks/misc/useOrgOptedIntoAi'
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils'
|
||||
import { getOrgAIDetails } from 'lib/ai/org-ai-details'
|
||||
import { RLS_PROMPT } from 'lib/ai/prompts'
|
||||
import { getTools } from 'lib/ai/tools'
|
||||
@@ -90,13 +91,9 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
}
|
||||
|
||||
try {
|
||||
const {
|
||||
model,
|
||||
error: modelError,
|
||||
providerOptions,
|
||||
} = await getModel({
|
||||
const { modelParams, error: modelError } = await getModel({
|
||||
provider: 'openai',
|
||||
routingKey: 'sql-policy',
|
||||
modelEntry: DEFAULT_COMPLETION_MODEL,
|
||||
})
|
||||
|
||||
if (modelError) {
|
||||
@@ -112,8 +109,7 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
})
|
||||
|
||||
const { experimental_output } = await generateText({
|
||||
model,
|
||||
providerOptions,
|
||||
...modelParams,
|
||||
stopWhen: stepCountIs(5),
|
||||
prompt: source`
|
||||
You are a Postgres RLS (Row Level Security) expert.
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { generateText, Output } from 'ai'
|
||||
import { source } from 'common-tags'
|
||||
import { getModel } from 'lib/ai/model'
|
||||
import { DEFAULT_COMPLETION_MODEL } from 'lib/ai/model.utils'
|
||||
import apiWrapper from 'lib/api/apiWrapper'
|
||||
import { NextApiRequest, NextApiResponse } from 'next'
|
||||
import { z } from 'zod'
|
||||
@@ -38,13 +39,9 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
}
|
||||
|
||||
try {
|
||||
const {
|
||||
model,
|
||||
error: modelError,
|
||||
providerOptions,
|
||||
} = await getModel({
|
||||
const { modelParams, error: modelError } = await getModel({
|
||||
provider: 'openai',
|
||||
routingKey: 'sql',
|
||||
modelEntry: DEFAULT_COMPLETION_MODEL,
|
||||
})
|
||||
|
||||
if (modelError) {
|
||||
@@ -52,8 +49,7 @@ export async function handlePost(req: NextApiRequest, res: NextApiResponse) {
|
||||
}
|
||||
|
||||
const result = await generateText({
|
||||
model,
|
||||
providerOptions,
|
||||
...modelParams,
|
||||
output: Output.object({ schema: titleSchema }),
|
||||
prompt: source`
|
||||
Generate a short title and summarized description for this Postgres SQL snippet:
|
||||
|
||||
@@ -8,6 +8,7 @@ import { proxy, ref, snapshot, subscribe, useSnapshot } from 'valtio'
|
||||
|
||||
import { constructHeaders } from 'data/fetchers'
|
||||
import { prepareMessagesForAPI } from 'lib/ai/message-utils'
|
||||
import type { AssistantModelId } from 'lib/ai/model.utils'
|
||||
import { BASE_PATH, IS_PLATFORM } from 'lib/constants'
|
||||
|
||||
import { LOCAL_STORAGE_KEYS } from 'common'
|
||||
@@ -22,7 +23,7 @@ export type AssistantMessageType = MessageType
|
||||
|
||||
export type SqlSnippet = string | { label: string; content: string }
|
||||
|
||||
export type AssistantModel = 'gpt-5' | 'gpt-5-mini'
|
||||
export type AssistantModel = AssistantModelId
|
||||
|
||||
type ChatSession = {
|
||||
id: string
|
||||
@@ -287,7 +288,7 @@ function createChatInstance(
|
||||
const messages = chatInstance.messages
|
||||
const chat = state.chats[options.id]
|
||||
if (chat) {
|
||||
chat.messages = messages as AssistantMessageType[]
|
||||
chat.messages = messages
|
||||
chat.updatedAt = new Date()
|
||||
}
|
||||
|
||||
@@ -452,7 +453,7 @@ export const createAiAssistantState = (): AiAssistantState => {
|
||||
if (index !== -1) {
|
||||
state.updateMessage(msg)
|
||||
} else {
|
||||
messagesToAdd.push(msg as AssistantMessageType)
|
||||
messagesToAdd.push(msg)
|
||||
}
|
||||
})
|
||||
|
||||
@@ -468,7 +469,7 @@ export const createAiAssistantState = (): AiAssistantState => {
|
||||
|
||||
const messageIndex = chat.messages.findIndex((msg) => msg.id === updatedMessage.id)
|
||||
if (messageIndex !== -1) {
|
||||
chat.messages[messageIndex] = updatedMessage as AssistantMessageType
|
||||
chat.messages[messageIndex] = updatedMessage
|
||||
chat.updatedAt = new Date()
|
||||
}
|
||||
},
|
||||
|
||||
Reference in new issue
Block a user