From a0b73cdcecbc59b6534fbb1df4061f829497bff5 Mon Sep 17 00:00:00 2001 From: DarkSky <25152247+darkskygit@users.noreply.github.com> Date: Thu, 18 Sep 2025 18:51:12 +0800 Subject: [PATCH] feat: improve model resolve (#13601) fix AI-419 --- .docker/selfhost/schema.json | 4 +- .../__snapshots__/copilot.spec.ts.md | 34 ++++++ .../__snapshots__/copilot.spec.ts.snap | Bin 2366 -> 2608 bytes .../server/src/__tests__/copilot.spec.ts | 101 ++++++++++++++++++ .../server/src/plugins/copilot/config.ts | 2 +- .../server/src/plugins/copilot/controller.ts | 10 +- .../src/plugins/copilot/prompt/prompts.ts | 9 +- .../src/plugins/copilot/providers/openai.ts | 72 ++++++++++--- .../src/plugins/copilot/providers/types.ts | 1 + .../server/src/plugins/copilot/session.ts | 66 +++++++++++- .../server/src/plugins/payment/service.ts | 2 +- 11 files changed, 271 insertions(+), 30 deletions(-) diff --git a/.docker/selfhost/schema.json b/.docker/selfhost/schema.json index 097836224..d8356b1dc 100644 --- a/.docker/selfhost/schema.json +++ b/.docker/selfhost/schema.json @@ -669,12 +669,12 @@ }, "scenarios": { "type": "object", - "description": "Use custom models in scenarios and override default settings.\n@default {\"override_enabled\":false,\"scenarios\":{\"audio_transcribing\":\"gemini-2.5-flash\",\"chat\":\"claude-sonnet-4@20250514\",\"embedding\":\"gemini-embedding-001\",\"image\":\"gpt-image-1\",\"rerank\":\"gpt-4.1\",\"coding\":\"claude-sonnet-4@20250514\",\"complex_text_generation\":\"gpt-4o-2024-08-06\",\"quick_decision_making\":\"gpt-5-mini\",\"quick_text_generation\":\"gemini-2.5-flash\",\"polish_and_summarize\":\"gemini-2.5-flash\"}}", + "description": "Use custom models in scenarios and override default settings.\n@default {\"override_enabled\":false,\"scenarios\":{\"audio_transcribing\":\"gemini-2.5-flash\",\"chat\":\"gemini-2.5-flash\",\"embedding\":\"gemini-embedding-001\",\"image\":\"gpt-image-1\",\"rerank\":\"gpt-4.1\",\"coding\":\"claude-sonnet-4@20250514\",\"complex_text_generation\":\"gpt-4o-2024-08-06\",\"quick_decision_making\":\"gpt-5-mini\",\"quick_text_generation\":\"gemini-2.5-flash\",\"polish_and_summarize\":\"gemini-2.5-flash\"}}", "default": { "override_enabled": false, "scenarios": { "audio_transcribing": "gemini-2.5-flash", - "chat": "claude-sonnet-4@20250514", + "chat": "gemini-2.5-flash", "embedding": "gemini-embedding-001", "image": "gpt-image-1", "rerank": "gpt-4.1", diff --git a/packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.md b/packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.md index 1809ad50c..0797e8989 100644 --- a/packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.md +++ b/packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.md @@ -444,3 +444,37 @@ Generated by [AVA](https://avajs.dev). }, ], } + +## should resolve model correctly based on subscription status and prompt config + +> should honor requested pro model + + 'gemini-2.5-pro' + +> should fallback to default model + + 'gemini-2.5-flash' + +> should fallback to default model when requesting pro model during trialing + + 'gemini-2.5-flash' + +> should honor requested non-pro model during trialing + + 'gemini-2.5-flash' + +> should pick default model when no requested model during trialing + + 'gemini-2.5-flash' + +> should pick first pro model when no requested model during active + + 'gemini-2.5-pro' + +> should honor requested pro model during active + + 'claude-sonnet-4@20250514' + +> should fallback to default model when requesting non-optional model during active + + 'gemini-2.5-flash' diff --git a/packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.snap b/packages/backend/server/src/__tests__/__snapshots__/copilot.spec.ts.snap index 11157ea0b6d7d4fab32eb5c06dc831074f367862..cd677a05b9ee4e665100f8d615931f0fc8ed06c9 100644 GIT binary patch literal 2608 zcmV-03eWXHRzVr z^GH(KyzmOKzQ8K(usnAb-lBtD;x6}?;-)HdvuL}_F)Lh3R^q+rUc1qqN=CO}7mM8E zu8N{-&fBVNGFPZF_sr1rhzM?JC4!>xW|yIXQyM7ru?C=ICIRdPZ~(yJG_V)}MDB~h z@~{TW5FtPGKhi+tO91{3;06NhOnR#YH<_aZI6;7q5#Y-y*vXR24+!v60{oExuO+AV zUpmxGx@xGU>58B>Ldbi&)1o1`5*>9*cy~J1F)3}Sn5zP;KSY6tX=kiwn@9UPjt*Q2 z7VLq$2k!q{aA(FEz}>$Oz#RbY0&pBaehHbCG?4cNcwUG>{+&9=pVJ^8BBYNX%7;n+ zVXDRR$1uiD1 zXIrNKmjW3Bc$)#-W+bK$UV372<4Ppf=eyIgs>C|iosP!DT7e>iF&3zKdaX<`^7(wA z;l)@D?9yYG9=lwN?4q01oAgYN&vCEb-de)*xTT!M zdZYR&O-RO`wIrn9)CuW7hzV(Xosgo^-EaEOq!NBF>zA-?7VG4{VnI#!x zVjQEueH8d`DrXrHXI)56ZLLkaCZIpBzHdl7e<8b8mYJz?Cb}S!RAIX>)1oH3Acau_ zcw=(+xRleS?K&bSZo@_7cWXqDwjRd%f00);k-h6t7ohHEN+EoXw8TKmz{FVT(Ca1P$*d_}2 zZSnq;3|mdw;ESiHGZ#2h!{Rma;5D=O|K(7t0Wj*7ltKPvh? zfENIa5@35eyzeBy5dwJW@ctqJzD9r_rNDbNnEk`r9kpM1%w`FVHrv+~YrRsYSIS(? zoZKsAdZkRSl3 z*bHH~FTT`L(sk15l*=fg-$w{}6sZ7Ei`CHGnGovqX#l@YDroEE5kmgdpum{ri@`G2 z`$y7CCDb%&wlCA>?#+6CF|ml7%Su({-W@ym0;^OVK5mK14#lO)1>PN>sh0F3ghbD5 zQLMCU%v7v25&Ya8{G1GaPK}v5(wmd#bYZxp@;IwREOO|wvB)RtEb{G`Md*_Nevx{S zh`tKowS-2aNn_V9Uk56pz@7KUSuj|4Y6H?TBr<+C0d6J0JOLiL^!rRHF$$n`zB{t2 z??BHJ;9@#^hJgKNgO-e9B^e-g^HRNiGX>s7o4wwjidbL9WD6$!-%e@dy(8urx4^yN ze7Nd~N>zn7n?g%X=4*)|u zoq|q>0KfOgjrjI&$NALalgg7B3yWG5sGtaf%(rwzCxt8r?lzciGS9D%b$M{_ z`72^K9bLWMBh&s=+bx+vbhWX2w6U&GZ4K5}+)UJh7C) zkP1J*Dsd`|9<5a>%v&@gFl+yfYyZ_`OnwoO7KNuhsU2$zH{#N_H*={as z)W+-0+w)c(?{Q6)g@WytLdm;~Sn^>lIh)dF(m4b8 zNb*Fg^E>sHdmi*>2JmE3NavU9>sw)(0fv+B{)bab*uEK&uNSae3(R&FBj4q;$8y|{ zYfi?2O_j5)SxywqTqO)=PU(223L0y_AQi?iR}fZ?$5GUjwkdeVry)~XWnQQ`K8ae) zaipgFCoZ*Tq!s`>y72FoXC+b&G5?^@MS- z-S&3Ring~ALfnR>XJ!JrS)V_9PS1i;S${4!x4P0$d19>B%Kq;>88 zmgJ?F^hZdTr*T-Oq6Uk8u-!Rn@&XK@3;ODG3iJDS#=gE%=0D$o-GA3%jM6sXFt$)& zD+TVMotphDF-}sz(nQ{Y@jDZIOaJxe<@h)t;=t)=XdH#rJbr;je-K3fLV;IOU-ZyT z24E&xGMh?p>+~AxoRZbQ+12Izi7@Tvk~NXeH(ePWiW^fYoH=d=snO6^=9%OLQ@EzA z<)!7>)j)()Ox2{x+(J}+S;8&arF!-B8O>dr<9zPJMlw?372CCQ6XR34swXrtm{D~m zT;ZRyJy_#X@j!IA9xZE2ldRDOSw+Vu*k)-(<~Vs~&H6`k1zuz|N3BBI!8B=M-L4ny zGk3gdvrzN=UntLJKBsr3cXAn@xWdg{9u51ohN^ANwlm2U4Qe~7TS literal 2366 zcmV-E3BmS3RzVybUkh98Rv00000000B+TWf3_R~7!wow2>KW5=Nt&?vM+n}~pHr+GmdMU<9UK--YE ziIk=#mF~`7?~XH%W$vsWJOxcE6_Bb5#7p@>BVJ1TBS=UfUX@5IRS2jAYAb$3AW&QJ zs*0%iQ3S#>b7yvace9gq>NG*?AM07)Gv~~?=R4m$_s;R7Gi5ihkI4s5iO7juv9MXz zf`(hM%EIDhS6ESCdEDn!Vc9hvSyD*p1imAEUYE5XlAk?Qsqly&p~npjA^xEf88-rW zH-Ps6xQYhX)4(lZ$QOySt99$vfrdd0G(vQHdgBPu$Ku}<>oC6s;4uKt0eCsnGkBmp zE9{6^emOoqo|s%uz*`8|OPZ*;Tyzv?QTBbLCR{hL=7Z3! zfX0Z95|9lJ)_KT1si`o-z!kZOT$GXU5?idm)<9ybn>Mo{MHsutjuY@X-94ktJ<_$! z2Vo_*mv|*yUKZ}c+f}fu!WSWrgca3HOX`W>yGu(^3j2=qeGX11^eLc zgZuv$+~!0JxCi$FxD~)101g8vFCo(_0{K9K=Y<^P-)VyUIR)|&q5(pd4?E*aGXbYf zfU~p^e~Ey{i=qDv0naKNze%6#cIfY5z)cLep8<2a_+Xp(Nd}x^z*7u3r;C@`#Q)0x z(*WLK05=<&`0&LimM|_wVtt`ED=SK@)4kbfO{@zrWN^+CGf%IUD@M6oPAt5*az~a* zJ<|k+&DCQFZFEi7(UCs8^x37)E`4^n64^y|RYm`c0FN{)3#5wloPTE8XKw|H(o^GM*sdDFR+Btn!T^28=RbN@o^5#$+5~z`YFk zcp+yQ8E2ipy203}MeYKfv zXPe3PvP!lRnb#0-9Ras3?Rd+@8Fqqz&v&X0+B57|1pJnOS9Q_$4BNnfcw2l+L58g) zZOZ+$>3EJBZ@nKb*nanf^oe|N@q#&Fid#ah3V>6bE>Gbi^;nSLqL zFJ=0rOuv-5Y%%-Bmon-gXH$FELD=vu$BSyfOEy{ znF98oEmq8kS7n0OjZ5|R^$d6qTR!W9g@}!H>}Y9xB+c<2NTY1x5aaUv)1Ws`zVqyWF? z$F2B|yW?VNiAm+S4oWHv@zrsY)nO;9-I+M^e(t&=h1bxiCORk<92v>vD_f@x4kRvT z->x|vw{AJo+RyEp;|uE!!9(Bit2e9JoNPO9PiAAjCj1uHEgr`8u__PlIde(;X5%aO zdu%!ub^NN8L{~YxTRH0oQAdYWi<`yWQ1Hs4RpT-pR(^lu^7h*p)w%%B=cgZ7ake*i zJc)rT%*j2)PPWw_%RT*9YlCZCJjC0&v&CK()nLc&i)Z)f9&37`?sCVsKHA`Z)bK1j z@al~OmF3V8Gw$N>efc9E-DBO=Cf1Nb=68F-u5sV7rBx3d-*)P*kW<>D45w81VeaLr zFn+M%c|2USGBB(8n^N=FkU9BfL^>0MY9`ghI)R^Y={uKmsU6fEH;61d418-gD9cVS z;2y|20uHZ6LGayXLGTQKXNwDhzXJG&F4lHx#D)kM)fd9{x9{mfcth(#$isLc%*XER zGWN;#u@af535XZ)7YKNTx|S|v{+p5sNLJfhwH0?r7v9f+gA90t0Vg%#kq)!Ie0koa z;ytCPa>{lE_Z#(TuO2O`^;D+Ei7XFL-1HC!sFT9Yy231#_d8NLel?-*RO8>88i$SZ zYn;q9&IjQ!S?9LM9`aDgWldD}1@^R87L|(QS5wW~j9l{prMXnF&t!83@Cp4ys{1?j zwmuK`GXr>1SJM6E`r3Awn!u?3?tiqfgdLa-%FP0HQ-wS3Vm5b^YQ`o-9M?&i2ewep z=H+^_Xij=*I42_&&q75LogXMlV>np}>`9SFQBm54it?Yh z*gYe&0XW6`#WcA;jSk7pz6s!T*ZoLpSF4Qo@hIc&jcPiqBtBW)B3oTOV%*kgdAs|9 zmbWn?zva9l&8+%YaP`xSMs` z?3c(m$^cstc^Ag-PVjB}ueY6#j}szJCjBJKqp*_O&$8?vg4kae@Jiu}9=5>%ES)7w kg#@=oJwuaIO6o7$dX(Rurrn%g6Q%j~KZ?GO%TPA}0Otgq%m4rY diff --git a/packages/backend/server/src/__tests__/copilot.spec.ts b/packages/backend/server/src/__tests__/copilot.spec.ts index 978cee77f..6a3b69e75 100644 --- a/packages/backend/server/src/__tests__/copilot.spec.ts +++ b/packages/backend/server/src/__tests__/copilot.spec.ts @@ -60,6 +60,9 @@ import { import { AutoRegisteredWorkflowExecutor } from '../plugins/copilot/workflow/executor/utils'; import { WorkflowGraphList } from '../plugins/copilot/workflow/graph'; import { CopilotWorkspaceService } from '../plugins/copilot/workspace'; +import { PaymentModule } from '../plugins/payment'; +import { SubscriptionService } from '../plugins/payment/service'; +import { SubscriptionStatus } from '../plugins/payment/types'; import { MockCopilotProvider } from './mocks'; import { createTestingModule, TestingModule } from './utils'; import { WorkflowTestCases } from './utils/copilot'; @@ -82,6 +85,7 @@ type Context = { storage: CopilotStorage; workflow: CopilotWorkflowService; cronJobs: CopilotCronJobs; + subscription: SubscriptionService; executors: { image: CopilotChatImageExecutor; text: CopilotChatTextExecutor; @@ -116,6 +120,7 @@ test.before(async t => { }, }, }), + PaymentModule, QuotaModule, StorageModule, CopilotModule, @@ -124,6 +129,13 @@ test.before(async t => { // use real JobQueue for testing builder.overrideProvider(JobQueue).useClass(JobQueue); builder.overrideProvider(OpenAIProvider).useClass(MockCopilotProvider); + builder.overrideProvider(SubscriptionService).useClass( + class { + select() { + return { getSubscription: async () => undefined }; + } + } + ); }, }); @@ -145,6 +157,7 @@ test.before(async t => { const transcript = module.get(CopilotTranscriptionService); const workspaceEmbedding = module.get(CopilotWorkspaceService); const cronJobs = module.get(CopilotCronJobs); + const subscription = module.get(SubscriptionService); t.context.module = module; t.context.auth = auth; @@ -163,6 +176,7 @@ test.before(async t => { t.context.transcript = transcript; t.context.workspaceEmbedding = workspaceEmbedding; t.context.cronJobs = cronJobs; + t.context.subscription = subscription; t.context.executors = { image: module.get(CopilotChatImageExecutor), @@ -2047,3 +2061,90 @@ test('should handle copilot cron jobs correctly', async t => { toBeGenerateStub.restore(); jobAddStub.restore(); }); + +test('should resolve model correctly based on subscription status and prompt config', async t => { + const { db, session, subscription } = t.context; + + // 1) Seed a prompt that has optionalModels and proModels in config + const promptName = 'resolve-model-test'; + await db.aiPrompt.create({ + data: { + name: promptName, + model: 'gemini-2.5-flash', + messages: { + create: [{ idx: 0, role: 'system', content: 'test' }], + }, + config: { proModels: ['gemini-2.5-pro', 'claude-sonnet-4@20250514'] }, + optionalModels: [ + 'gemini-2.5-flash', + 'gemini-2.5-pro', + 'claude-sonnet-4@20250514', + ], + }, + }); + + // 2) Create a chat session with this prompt + const sessionId = await session.create({ + promptName, + docId: 'test', + workspaceId: 'test', + userId, + pinned: false, + }); + const s = (await session.get(sessionId))!; + + const mockStatus = (status?: SubscriptionStatus) => { + Sinon.restore(); + Sinon.stub(subscription, 'select').callsFake(() => ({ + // @ts-expect-error mock + getSubscription: async () => (status ? { status } : null), + })); + }; + + // payment disabled -> allow requested if in optional; pro not blocked + { + const model1 = await s.resolveModel(false, 'gemini-2.5-pro'); + t.snapshot(model1, 'should honor requested pro model'); + + const model2 = await s.resolveModel(false, 'not-in-optional'); + t.snapshot(model2, 'should fallback to default model'); + } + + // payment enabled + trialing: requesting pro should fallback to default + { + mockStatus(SubscriptionStatus.Trialing); + const model3 = await s.resolveModel(true, 'gemini-2.5-pro'); + t.snapshot( + model3, + 'should fallback to default model when requesting pro model during trialing' + ); + + const model4 = await s.resolveModel(true, 'gemini-2.5-flash'); + t.snapshot(model4, 'should honor requested non-pro model during trialing'); + + const model5 = await s.resolveModel(true); + t.snapshot( + model5, + 'should pick default model when no requested model during trialing' + ); + } + + // payment enabled + active: without requested -> first pro; requested pro should be honored + { + mockStatus(SubscriptionStatus.Active); + const model6 = await s.resolveModel(true); + t.snapshot( + model6, + 'should pick first pro model when no requested model during active' + ); + + const model7 = await s.resolveModel(true, 'claude-sonnet-4@20250514'); + t.snapshot(model7, 'should honor requested pro model during active'); + + const model8 = await s.resolveModel(true, 'not-in-optional'); + t.snapshot( + model8, + 'should fallback to default model when requesting non-optional model during active' + ); + } +}); diff --git a/packages/backend/server/src/plugins/copilot/config.ts b/packages/backend/server/src/plugins/copilot/config.ts index 813e74601..bc025f661 100644 --- a/packages/backend/server/src/plugins/copilot/config.ts +++ b/packages/backend/server/src/plugins/copilot/config.ts @@ -51,7 +51,7 @@ defineModuleConfig('copilot', { override_enabled: false, scenarios: { audio_transcribing: 'gemini-2.5-flash', - chat: 'claude-sonnet-4@20250514', + chat: 'gemini-2.5-flash', embedding: 'gemini-embedding-001', image: 'gpt-image-1', rerank: 'gpt-4.1', diff --git a/packages/backend/server/src/plugins/copilot/controller.ts b/packages/backend/server/src/plugins/copilot/controller.ts index 240a1051a..c54f76ecc 100644 --- a/packages/backend/server/src/plugins/copilot/controller.ts +++ b/packages/backend/server/src/plugins/copilot/controller.ts @@ -44,6 +44,7 @@ import { NoCopilotProviderAvailable, UnsplashIsNotConfigured, } from '../../base'; +import { ServerFeature, ServerService } from '../../core'; import { CurrentUser, Public } from '../../core/auth'; import { CopilotContextService } from './context'; import { @@ -75,6 +76,7 @@ export class CopilotController implements BeforeApplicationShutdown { constructor( private readonly config: Config, + private readonly server: ServerService, private readonly chatSession: ChatSessionService, private readonly context: CopilotContextService, private readonly provider: CopilotProviderFactory, @@ -112,10 +114,10 @@ export class CopilotController implements BeforeApplicationShutdown { throw new CopilotSessionNotFound(); } - const model = - modelId && session.optionalModels.includes(modelId) - ? modelId - : session.model; + const model = await session.resolveModel( + this.server.features.includes(ServerFeature.Payment), + modelId + ); const hasAttachment = messageId ? !!(await session.getMessageById(messageId)).attachments?.length diff --git a/packages/backend/server/src/plugins/copilot/prompt/prompts.ts b/packages/backend/server/src/plugins/copilot/prompt/prompts.ts index 1b53ea983..07e9d1ad5 100644 --- a/packages/backend/server/src/plugins/copilot/prompt/prompts.ts +++ b/packages/backend/server/src/plugins/copilot/prompt/prompts.ts @@ -1928,7 +1928,7 @@ Now apply the \`updates\` to the \`content\`, following the intent in \`op\`, an ]; const CHAT_PROMPT: Omit = { - model: 'claude-sonnet-4@20250514', + model: 'gemini-2.5-flash', optionalModels: [ 'gpt-4.1', 'gpt-5', @@ -2099,6 +2099,13 @@ Below is the user's query. Please respond in the user's preferred language witho 'codeArtifact', 'blobRead', ], + proModels: [ + 'gemini-2.5-pro', + 'claude-opus-4@20250514', + 'claude-sonnet-4@20250514', + 'claude-3-7-sonnet@20250219', + 'claude-3-5-sonnet-v2@20241022', + ], }, }; diff --git a/packages/backend/server/src/plugins/copilot/providers/openai.ts b/packages/backend/server/src/plugins/copilot/providers/openai.ts index 12160a1ad..ba872bab8 100644 --- a/packages/backend/server/src/plugins/copilot/providers/openai.ts +++ b/packages/backend/server/src/plugins/copilot/providers/openai.ts @@ -4,6 +4,10 @@ import { type OpenAIProvider as VercelOpenAIProvider, OpenAIResponsesProviderOptions, } from '@ai-sdk/openai'; +import { + createOpenAICompatible, + type OpenAICompatibleProvider as VercelOpenAICompatibleProvider, +} from '@ai-sdk/openai-compatible'; import { AISDKError, embedMany, @@ -18,6 +22,7 @@ import { z } from 'zod'; import { CopilotPromptInvalid, + CopilotProviderNotSupported, CopilotProviderSideError, metrics, UserFriendlyError, @@ -47,6 +52,7 @@ export const DEFAULT_DIMENSIONS = 256; export type OpenAIConfig = { apiKey: string; baseURL?: string; + oldApiStyle?: boolean; }; const ModelListSchema = z.object({ @@ -296,7 +302,7 @@ export class OpenAIProvider extends CopilotProvider { }, ]; - #instance!: VercelOpenAIProvider; + #instance!: VercelOpenAIProvider | VercelOpenAICompatibleProvider; override configured(): boolean { return !!this.config.apiKey; @@ -304,10 +310,17 @@ export class OpenAIProvider extends CopilotProvider { protected override setup() { super.setup(); - this.#instance = createOpenAI({ - apiKey: this.config.apiKey, - baseURL: this.config.baseURL, - }); + this.#instance = + this.config.oldApiStyle && this.config.baseURL + ? createOpenAICompatible({ + name: 'openai-compatible-old-style', + apiKey: this.config.apiKey, + baseURL: this.config.baseURL, + }) + : createOpenAI({ + apiKey: this.config.apiKey, + baseURL: this.config.baseURL, + }); } private handleError( @@ -341,7 +354,7 @@ export class OpenAIProvider extends CopilotProvider { override async refreshOnlineModels() { try { const baseUrl = this.config.baseURL || 'https://api.openai.com/v1'; - if (baseUrl && !this.onlineModelList.length) { + if (this.config.apiKey && baseUrl && !this.onlineModelList.length) { const { data } = await fetch(`${baseUrl}/models`, { headers: { Authorization: `Bearer ${this.config.apiKey}`, @@ -361,7 +374,11 @@ export class OpenAIProvider extends CopilotProvider { toolName: CopilotChatTools, model: string ): [string, Tool?] | undefined { - if (toolName === 'webSearch' && !this.isReasoningModel(model)) { + if ( + toolName === 'webSearch' && + 'responses' in this.#instance && + !this.isReasoningModel(model) + ) { return ['web_search_preview', openai.tools.webSearchPreview({})]; } else if (toolName === 'docEdit') { return ['doc_edit', undefined]; @@ -374,10 +391,7 @@ export class OpenAIProvider extends CopilotProvider { messages: PromptMessage[], options: CopilotChatOptions = {} ): Promise { - const fullCond = { - ...cond, - outputType: ModelOutputType.Text, - }; + const fullCond = { ...cond, outputType: ModelOutputType.Text }; await this.checkParams({ messages, cond: fullCond, options }); const model = this.selectModel(fullCond); @@ -386,7 +400,10 @@ export class OpenAIProvider extends CopilotProvider { const [system, msgs] = await chatToGPTMessage(messages); - const modelInstance = this.#instance.responses(model.id); + const modelInstance = + 'responses' in this.#instance + ? this.#instance.responses(model.id) + : this.#instance(model.id); const { text } = await generateText({ model: modelInstance, @@ -507,7 +524,10 @@ export class OpenAIProvider extends CopilotProvider { throw new CopilotPromptInvalid('Schema is required'); } - const modelInstance = this.#instance.responses(model.id); + const modelInstance = + 'responses' in this.#instance + ? this.#instance.responses(model.id) + : this.#instance(model.id); const { object } = await generateObject({ model: modelInstance, @@ -539,7 +559,10 @@ export class OpenAIProvider extends CopilotProvider { await this.checkParams({ messages: [], cond: fullCond, options }); const model = this.selectModel(fullCond); // get the log probability of "yes"/"no" - const instance = this.#instance.chat(model.id); + const instance = + 'chat' in this.#instance + ? this.#instance.chat(model.id) + : this.#instance(model.id); const scores = await Promise.all( chunkMessages.map(async messages => { @@ -600,7 +623,10 @@ export class OpenAIProvider extends CopilotProvider { options: CopilotChatOptions = {} ) { const [system, msgs] = await chatToGPTMessage(messages); - const modelInstance = this.#instance.responses(model.id); + const modelInstance = + 'responses' in this.#instance + ? this.#instance.responses(model.id) + : this.#instance(model.id); const { fullStream } = streamText({ model: modelInstance, system, @@ -685,6 +711,13 @@ export class OpenAIProvider extends CopilotProvider { await this.checkParams({ messages, cond: fullCond, options }); const model = this.selectModel(fullCond); + if (!('image' in this.#instance)) { + throw new CopilotProviderNotSupported({ + provider: this.type, + kind: 'image', + }); + } + metrics.ai .counter('generate_images_stream_calls') .add(1, { model: model.id }); @@ -735,6 +768,13 @@ export class OpenAIProvider extends CopilotProvider { await this.checkParams({ embeddings: messages, cond: fullCond, options }); const model = this.selectModel(fullCond); + if (!('embedding' in this.#instance)) { + throw new CopilotProviderNotSupported({ + provider: this.type, + kind: 'embedding', + }); + } + try { metrics.ai .counter('generate_embedding_calls') @@ -775,6 +815,6 @@ export class OpenAIProvider extends CopilotProvider { private isReasoningModel(model: string) { // o series reasoning models - return model.startsWith('o'); + return model.startsWith('o') || model.startsWith('gpt-5'); } } diff --git a/packages/backend/server/src/plugins/copilot/providers/types.ts b/packages/backend/server/src/plugins/copilot/providers/types.ts index fb4cb2ae9..e568be80e 100644 --- a/packages/backend/server/src/plugins/copilot/providers/types.ts +++ b/packages/backend/server/src/plugins/copilot/providers/types.ts @@ -80,6 +80,7 @@ export const PromptToolsSchema = z export const PromptConfigStrictSchema = z.object({ tools: PromptToolsSchema.nullable().optional(), + proModels: z.array(z.string()).nullable().optional(), // params requirements requireContent: z.boolean().nullable().optional(), requireAttachment: z.boolean().nullable().optional(), diff --git a/packages/backend/server/src/plugins/copilot/session.ts b/packages/backend/server/src/plugins/copilot/session.ts index a88b7435f..6ce10c5ad 100644 --- a/packages/backend/server/src/plugins/copilot/session.ts +++ b/packages/backend/server/src/plugins/copilot/session.ts @@ -25,6 +25,8 @@ import { type UpdateChatSession, UpdateChatSessionOptions, } from '../../models'; +import { SubscriptionService } from '../payment/service'; +import { SubscriptionPlan, SubscriptionStatus } from '../payment/types'; import { ChatMessageCache } from './message'; import { ChatPrompt, PromptService } from './prompt'; import { @@ -58,6 +60,7 @@ declare global { export class ChatSession implements AsyncDisposable { private stashMessageCount = 0; constructor( + private readonly moduleRef: ModuleRef, private readonly messageCache: ChatMessageCache, private readonly state: ChatSessionState, private readonly dispose?: (state: ChatSessionState) => Promise, @@ -72,6 +75,10 @@ export class ChatSession implements AsyncDisposable { return this.state.prompt.optionalModels; } + get proModels() { + return this.state.prompt.config?.proModels || []; + } + get config() { const { sessionId, @@ -93,6 +100,50 @@ export class ChatSession implements AsyncDisposable { return this.state.messages.findLast(m => m.role === 'user'); } + async resolveModel( + hasPayment: boolean, + requestedModelId?: string + ): Promise { + const defaultModel = this.model; + const normalize = (m?: string) => + !!m && this.optionalModels.includes(m) ? m : defaultModel; + const isPro = (m?: string) => !!m && this.proModels.includes(m); + + // try resolve payment subscription service lazily + let paymentEnabled = hasPayment; + let isUserAIPro = false; + try { + if (paymentEnabled) { + const sub = this.moduleRef.get(SubscriptionService, { + strict: false, + }); + const subscription = await sub + .select(SubscriptionPlan.AI) + .getSubscription({ + userId: this.config.userId, + plan: SubscriptionPlan.AI, + } as any); + isUserAIPro = subscription?.status === SubscriptionStatus.Active; + } + } catch { + // payment not available -> skip checks + paymentEnabled = false; + } + + if (paymentEnabled) { + if (isUserAIPro) { + if (!requestedModelId) { + const firstPro = this.proModels[0]; + return normalize(firstPro); + } + } else if (isPro(requestedModelId)) { + return defaultModel; + } + } + + return normalize(requestedModelId); + } + push(message: ChatMessage) { if ( this.state.prompt.action && @@ -539,12 +590,17 @@ export class ChatSessionService { async get(sessionId: string): Promise { const state = await this.getSessionInfo(sessionId); if (state) { - return new ChatSession(this.messageCache, state, async state => { - await this.models.copilotSession.updateMessages(state); - if (!state.prompt.action) { - await this.jobs.add('copilot.session.generateTitle', { sessionId }); + return new ChatSession( + this.moduleRef, + this.messageCache, + state, + async state => { + await this.models.copilotSession.updateMessages(state); + if (!state.prompt.action) { + await this.jobs.add('copilot.session.generateTitle', { sessionId }); + } } - }); + ); } return null; } diff --git a/packages/backend/server/src/plugins/payment/service.ts b/packages/backend/server/src/plugins/payment/service.ts index a4d13d087..f08e07c7f 100644 --- a/packages/backend/server/src/plugins/payment/service.ts +++ b/packages/backend/server/src/plugins/payment/service.ts @@ -89,7 +89,7 @@ export class SubscriptionService { return this.stripeProvider.stripe; } - private select(plan: SubscriptionPlan): SubscriptionManager { + select(plan: SubscriptionPlan): SubscriptionManager { switch (plan) { case SubscriptionPlan.Team: return this.workspaceManager;