add some parallelization to ai gateway pipeline

This commit is contained in:
miloschwartz
2026-08-07 14:21:13 -04:00
parent 12056aebc6
commit bc7a883f6c
+125 -113
View File
@@ -242,137 +242,146 @@ async function resolveRequestUser(
} }
async function resolveTarget(host: string): Promise<ResolvedTarget | null> { async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
const [resourceRow] = await db const [[resourceRow], [siteResourceRow]] = await Promise.all([
.select({ db
resourceId: resources.resourceId,
orgId: resources.orgId
})
.from(resources)
.where(
and(
eq(resources.fullDomain, host),
eq(resources.mode, "inference"),
eq(resources.enabled, true)
)
)
.limit(1);
if (resourceRow) {
const attachmentRows = await db
.select({ .select({
provider: aiProviders, resourceId: resources.resourceId,
accessMode: resourceAiProviders.accessMode orgId: resources.orgId
}) })
.from(resourceAiProviders) .from(resources)
.innerJoin(
aiProviders,
eq(resourceAiProviders.providerId, aiProviders.providerId)
)
.where( .where(
and( and(
eq(resourceAiProviders.resourceId, resourceRow.resourceId), eq(resources.fullDomain, host),
eq(aiProviders.enabled, true) eq(resources.mode, "inference"),
eq(resources.enabled, true)
) )
); )
.limit(1),
db
.select({
siteResourceId: siteResources.siteResourceId,
orgId: siteResources.orgId
})
.from(siteResources)
.where(
and(
eq(siteResources.fullDomain, host),
eq(siteResources.mode, "inference"),
eq(siteResources.enabled, true)
)
)
.limit(1)
]);
// Prefer public inference resources when both match the same host.
if (resourceRow) {
const [attachmentRows, resourcePatterns] = await Promise.all([
db
.select({
provider: aiProviders,
accessMode: resourceAiProviders.accessMode
})
.from(resourceAiProviders)
.innerJoin(
aiProviders,
eq(resourceAiProviders.providerId, aiProviders.providerId)
)
.where(
and(
eq(
resourceAiProviders.resourceId,
resourceRow.resourceId
),
eq(aiProviders.enabled, true)
)
),
db
.select({
providerId: aiModels.providerId,
modelKey: aiModels.modelKey,
listType: resourceAiModels.listType,
enabled: aiModels.enabled
})
.from(resourceAiModels)
.innerJoin(
aiModels,
eq(resourceAiModels.modelId, aiModels.modelId)
)
.where(eq(resourceAiModels.resourceId, resourceRow.resourceId))
]);
if (attachmentRows.length === 0) { if (attachmentRows.length === 0) {
return null; return null;
} }
const attachments: ProviderAttachment[] = attachmentRows.map((a) => ({
provider: a.provider,
accessMode: a.accessMode
}));
const resourcePatterns = await db
.select({
providerId: aiModels.providerId,
modelKey: aiModels.modelKey,
listType: resourceAiModels.listType,
enabled: aiModels.enabled
})
.from(resourceAiModels)
.innerJoin(aiModels, eq(resourceAiModels.modelId, aiModels.modelId))
.where(eq(resourceAiModels.resourceId, resourceRow.resourceId));
return { return {
resourceId: resourceRow.resourceId, resourceId: resourceRow.resourceId,
siteResourceId: null, siteResourceId: null,
orgId: resourceRow.orgId, orgId: resourceRow.orgId,
attachments, attachments: attachmentRows.map((a) => ({
provider: a.provider,
accessMode: a.accessMode
})),
resourceListsByProvider: groupPatternsByProvider(resourcePatterns) resourceListsByProvider: groupPatternsByProvider(resourcePatterns)
}; };
} }
const [siteResourceRow] = await db
.select({
siteResourceId: siteResources.siteResourceId,
orgId: siteResources.orgId
})
.from(siteResources)
.where(
and(
eq(siteResources.fullDomain, host),
eq(siteResources.mode, "inference"),
eq(siteResources.enabled, true)
)
)
.limit(1);
if (siteResourceRow) { if (siteResourceRow) {
const attachmentRows = await db const [attachmentRows, resourcePatterns] = await Promise.all([
.select({ db
provider: aiProviders, .select({
accessMode: siteResourceAiProviders.accessMode provider: aiProviders,
}) accessMode: siteResourceAiProviders.accessMode
.from(siteResourceAiProviders) })
.innerJoin( .from(siteResourceAiProviders)
aiProviders, .innerJoin(
eq(siteResourceAiProviders.providerId, aiProviders.providerId) aiProviders,
)
.where(
and(
eq( eq(
siteResourceAiProviders.siteResourceId, siteResourceAiProviders.providerId,
siteResourceRow.siteResourceId aiProviders.providerId
), )
eq(aiProviders.enabled, true)
) )
); .where(
and(
eq(
siteResourceAiProviders.siteResourceId,
siteResourceRow.siteResourceId
),
eq(aiProviders.enabled, true)
)
),
db
.select({
providerId: aiModels.providerId,
modelKey: aiModels.modelKey,
listType: siteResourceAiModels.listType,
enabled: aiModels.enabled
})
.from(siteResourceAiModels)
.innerJoin(
aiModels,
eq(siteResourceAiModels.modelId, aiModels.modelId)
)
.where(
eq(
siteResourceAiModels.siteResourceId,
siteResourceRow.siteResourceId
)
)
]);
if (attachmentRows.length === 0) { if (attachmentRows.length === 0) {
return null; return null;
} }
const attachments: ProviderAttachment[] = attachmentRows.map((a) => ({
provider: a.provider,
accessMode: a.accessMode
}));
const resourcePatterns = await db
.select({
providerId: aiModels.providerId,
modelKey: aiModels.modelKey,
listType: siteResourceAiModels.listType,
enabled: aiModels.enabled
})
.from(siteResourceAiModels)
.innerJoin(
aiModels,
eq(siteResourceAiModels.modelId, aiModels.modelId)
)
.where(
eq(
siteResourceAiModels.siteResourceId,
siteResourceRow.siteResourceId
)
);
return { return {
resourceId: null, resourceId: null,
siteResourceId: siteResourceRow.siteResourceId, siteResourceId: siteResourceRow.siteResourceId,
orgId: siteResourceRow.orgId, orgId: siteResourceRow.orgId,
attachments, attachments: attachmentRows.map((a) => ({
provider: a.provider,
accessMode: a.accessMode
})),
resourceListsByProvider: groupPatternsByProvider(resourcePatterns) resourceListsByProvider: groupPatternsByProvider(resourcePatterns)
}; };
} }
@@ -587,13 +596,6 @@ export async function handleAiGatewayProxy(
const { attachments, resourceListsByProvider, resourceId, orgId } = const { attachments, resourceListsByProvider, resourceId, orgId } =
target; target;
const requestUser = await resolveRequestUser(req, resourceId, orgId);
if (requestUser) {
logger.debug(
`AI gateway request from user ${requestUser.userId} (${requestUser.username})`
);
}
const capableAttachments = attachments.filter((a) => const capableAttachments = attachments.filter((a) =>
providerHasCapability(a.provider.capabilities, capability) providerHasCapability(a.provider.capabilities, capability)
); );
@@ -608,11 +610,21 @@ export async function handleAiGatewayProxy(
const requestedModel = def.extractModel(req); const requestedModel = def.extractModel(req);
const selection = await selectProvider( const [requestUser, selection] = await Promise.all([
capableAttachments, resolveRequestUser(req, resourceId, orgId),
resourceListsByProvider, selectProvider(
requestedModel capableAttachments,
); resourceListsByProvider,
requestedModel
)
]);
if (requestUser) {
logger.debug(
`AI gateway request from user ${requestUser.userId} (${requestUser.username})`
);
}
if (!selection.ok) { if (!selection.ok) {
return res.status(selection.status).json({ return res.status(selection.status).json({
error: { message: selection.message } error: { message: selection.message }