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
+88 -76
View File
@@ -242,7 +242,8 @@ async function resolveRequestUser(
}
async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
const [resourceRow] = await db
const [[resourceRow], [siteResourceRow]] = await Promise.all([
db
.select({
resourceId: resources.resourceId,
orgId: resources.orgId
@@ -255,56 +256,8 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
eq(resources.enabled, true)
)
)
.limit(1);
if (resourceRow) {
const attachmentRows = await 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)
)
);
if (attachmentRows.length === 0) {
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 {
resourceId: resourceRow.resourceId,
siteResourceId: null,
orgId: resourceRow.orgId,
attachments,
resourceListsByProvider: groupPatternsByProvider(resourcePatterns)
};
}
const [siteResourceRow] = await db
.limit(1),
db
.select({
siteResourceId: siteResources.siteResourceId,
orgId: siteResources.orgId
@@ -317,10 +270,65 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
eq(siteResources.enabled, true)
)
)
.limit(1);
.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) {
return null;
}
return {
resourceId: resourceRow.resourceId,
siteResourceId: null,
orgId: resourceRow.orgId,
attachments: attachmentRows.map((a) => ({
provider: a.provider,
accessMode: a.accessMode
})),
resourceListsByProvider: groupPatternsByProvider(resourcePatterns)
};
}
if (siteResourceRow) {
const attachmentRows = await db
const [attachmentRows, resourcePatterns] = await Promise.all([
db
.select({
provider: aiProviders,
accessMode: siteResourceAiProviders.accessMode
@@ -328,7 +336,10 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
.from(siteResourceAiProviders)
.innerJoin(
aiProviders,
eq(siteResourceAiProviders.providerId, aiProviders.providerId)
eq(
siteResourceAiProviders.providerId,
aiProviders.providerId
)
)
.where(
and(
@@ -338,18 +349,8 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
),
eq(aiProviders.enabled, true)
)
);
if (attachmentRows.length === 0) {
return null;
}
const attachments: ProviderAttachment[] = attachmentRows.map((a) => ({
provider: a.provider,
accessMode: a.accessMode
}));
const resourcePatterns = await db
),
db
.select({
providerId: aiModels.providerId,
modelKey: aiModels.modelKey,
@@ -366,13 +367,21 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
siteResourceAiModels.siteResourceId,
siteResourceRow.siteResourceId
)
);
)
]);
if (attachmentRows.length === 0) {
return null;
}
return {
resourceId: null,
siteResourceId: siteResourceRow.siteResourceId,
orgId: siteResourceRow.orgId,
attachments,
attachments: attachmentRows.map((a) => ({
provider: a.provider,
accessMode: a.accessMode
})),
resourceListsByProvider: groupPatternsByProvider(resourcePatterns)
};
}
@@ -587,13 +596,6 @@ export async function handleAiGatewayProxy(
const { attachments, resourceListsByProvider, resourceId, orgId } =
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) =>
providerHasCapability(a.provider.capabilities, capability)
);
@@ -608,11 +610,21 @@ export async function handleAiGatewayProxy(
const requestedModel = def.extractModel(req);
const selection = await selectProvider(
const [requestUser, selection] = await Promise.all([
resolveRequestUser(req, resourceId, orgId),
selectProvider(
capableAttachments,
resourceListsByProvider,
requestedModel
)
]);
if (requestUser) {
logger.debug(
`AI gateway request from user ${requestUser.userId} (${requestUser.username})`
);
}
if (!selection.ok) {
return res.status(selection.status).json({
error: { message: selection.message }