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> { async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
const [resourceRow] = await db const [[resourceRow], [siteResourceRow]] = await Promise.all([
db
.select({ .select({
resourceId: resources.resourceId, resourceId: resources.resourceId,
orgId: resources.orgId orgId: resources.orgId
@@ -255,56 +256,8 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
eq(resources.enabled, true) eq(resources.enabled, true)
) )
) )
.limit(1); .limit(1),
db
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
.select({ .select({
siteResourceId: siteResources.siteResourceId, siteResourceId: siteResources.siteResourceId,
orgId: siteResources.orgId orgId: siteResources.orgId
@@ -317,10 +270,65 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
eq(siteResources.enabled, true) 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) { if (siteResourceRow) {
const attachmentRows = await db const [attachmentRows, resourcePatterns] = await Promise.all([
db
.select({ .select({
provider: aiProviders, provider: aiProviders,
accessMode: siteResourceAiProviders.accessMode accessMode: siteResourceAiProviders.accessMode
@@ -328,7 +336,10 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
.from(siteResourceAiProviders) .from(siteResourceAiProviders)
.innerJoin( .innerJoin(
aiProviders, aiProviders,
eq(siteResourceAiProviders.providerId, aiProviders.providerId) eq(
siteResourceAiProviders.providerId,
aiProviders.providerId
)
) )
.where( .where(
and( and(
@@ -338,18 +349,8 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
), ),
eq(aiProviders.enabled, true) eq(aiProviders.enabled, true)
) )
); ),
db
if (attachmentRows.length === 0) {
return null;
}
const attachments: ProviderAttachment[] = attachmentRows.map((a) => ({
provider: a.provider,
accessMode: a.accessMode
}));
const resourcePatterns = await db
.select({ .select({
providerId: aiModels.providerId, providerId: aiModels.providerId,
modelKey: aiModels.modelKey, modelKey: aiModels.modelKey,
@@ -366,13 +367,21 @@ async function resolveTarget(host: string): Promise<ResolvedTarget | null> {
siteResourceAiModels.siteResourceId, siteResourceAiModels.siteResourceId,
siteResourceRow.siteResourceId siteResourceRow.siteResourceId
) )
); )
]);
if (attachmentRows.length === 0) {
return null;
}
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([
resolveRequestUser(req, resourceId, orgId),
selectProvider(
capableAttachments, capableAttachments,
resourceListsByProvider, resourceListsByProvider,
requestedModel 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 }