Files
easyai-ai-gateway/apps/web/src/lib/run-task.ts
T

210 lines
6.6 KiB
TypeScript

import type { GatewayTask } from '@easyai-ai-gateway/contracts';
import {
createCompatibleChatCompletion,
createEmbedding,
createImageEditTask,
createImageGenerationTask,
createRerank,
createResponse,
createVideoGenerationTask,
getAPITask,
} from '../api';
import type { TaskForm, TaskSubmissionMode } from '../types';
export interface RunTaskResponse {
localOnly?: boolean;
next?: Record<string, string>;
submissionMode: TaskSubmissionMode;
task: GatewayTask;
}
export interface RunTaskOptions {
requestBody?: Record<string, unknown>;
submissionMode?: TaskSubmissionMode;
}
const simulationParameterKeys = ['runMode', 'run_mode', 'simulation', 'testMode', 'test_mode'] as const;
export function applyTaskSubmissionMode(
input: Record<string, unknown>,
submissionMode: TaskSubmissionMode,
): Record<string, unknown> {
const body = { ...input };
for (const key of simulationParameterKeys) delete body[key];
if (submissionMode === 'simulation') {
body.runMode = 'simulation';
body.simulation = true;
}
return body;
}
export async function runTask(token: string, task: TaskForm, options: RunTaskOptions = {}): Promise<RunTaskResponse> {
const submissionMode = options.submissionMode ?? 'simulation';
const requestBody = task.kind === 'tasks.retrieve'
? { taskId: task.taskId }
: applyTaskSubmissionMode(options.requestBody ?? defaultRequestBody(task), submissionMode);
if (task.kind === 'chat.completions') {
const result = await createCompatibleChatCompletion(
token,
requestBody as Parameters<typeof createCompatibleChatCompletion>[1],
);
return { localOnly: true, submissionMode, task: compatibleTask(task, result, requestBody, submissionMode) };
}
if (task.kind === 'responses') {
const result = await createResponse(token, requestBody as Parameters<typeof createResponse>[1]);
return { localOnly: true, submissionMode, task: compatibleTask(task, result, requestBody, submissionMode) };
}
if (task.kind === 'embeddings') {
const result = await createEmbedding(token, requestBody as Parameters<typeof createEmbedding>[1]);
return { localOnly: true, submissionMode, task: compatibleTask(task, result, requestBody, submissionMode) };
}
if (task.kind === 'reranks') {
const result = await createRerank(token, requestBody as Parameters<typeof createRerank>[1]);
return { localOnly: true, submissionMode, task: compatibleTask(task, result, requestBody, submissionMode) };
}
if (task.kind === 'images.generations') {
const response = await createImageGenerationTask(
token,
requestBody as Parameters<typeof createImageGenerationTask>[1],
);
return { ...response, submissionMode };
}
if (task.kind === 'images.edits') {
const response = await createImageEditTask(token, requestBody as Parameters<typeof createImageEditTask>[1]);
return { ...response, submissionMode };
}
if (task.kind === 'videos.generations') {
const response = await createVideoGenerationTask(
token,
requestBody as unknown as Parameters<typeof createVideoGenerationTask>[1],
);
return { ...response, submissionMode };
}
if (task.kind === 'tasks.retrieve') {
const taskId = task.taskId?.trim();
if (!taskId) throw new Error('请输入要取回的 Task ID');
const result = await getAPITask(token, taskId);
return {
localOnly: true,
submissionMode: 'production',
task: compatibleTask(task, result as unknown as Record<string, unknown>, requestBody, 'production'),
};
}
throw new Error(`Unsupported task kind: ${task.kind}`);
}
function compatibleTask(
task: TaskForm,
result: Record<string, unknown>,
requestBody: Record<string, unknown>,
submissionMode: TaskSubmissionMode,
): GatewayTask {
const now = new Date().toISOString();
return {
id: `docs-${task.kind}-${Date.now()}`,
asyncMode: false,
createdAt: now,
finishedAt: now,
kind: task.kind,
model: typeof requestBody.model === 'string' ? requestBody.model : task.model,
modelType: modelTypeForKind(task.kind),
request: requestBody,
result,
runMode: submissionMode,
status: 'succeeded',
updatedAt: now,
userId: 'docs-runner',
};
}
function defaultRequestBody(task: TaskForm): Record<string, unknown> {
if (task.kind === 'chat.completions') {
return {
model: task.model,
messages: [{ role: 'user', content: task.prompt }],
stream: false,
};
}
if (task.kind === 'responses') {
return {
model: task.model,
input: task.prompt,
instructions: task.instructions,
previous_response_id: task.previousResponseId,
store: true,
stream: false,
};
}
if (task.kind === 'embeddings') {
return {
model: task.model,
input: embeddingInput(task.prompt),
dimensions: task.dimensions,
};
}
if (task.kind === 'reranks') {
return {
model: task.model,
query: task.prompt,
documents: rerankDocuments(task.documents),
top_n: task.topN,
};
}
if (task.kind === 'images.generations') {
return {
model: task.model,
prompt: task.prompt,
quality: 'medium',
size: '1024x1024',
};
}
if (task.kind === 'images.edits') {
return {
model: task.model,
prompt: task.prompt,
image: task.image,
mask: task.mask,
};
}
if (task.kind === 'videos.generations') {
return {
model: task.model,
content: [{ type: 'text', text: task.prompt }],
aspect_ratio: task.aspectRatio ?? '16:9',
resolution: task.resolution ?? '720p',
duration: task.duration ?? 5,
audio: task.outputAudio ?? true,
};
}
if (task.kind === 'tasks.retrieve') return { taskId: task.taskId };
return { model: task.model };
}
function embeddingInput(prompt: string) {
const lines = splitLines(prompt);
return lines.length > 1 ? lines : (lines[0] ?? prompt);
}
function rerankDocuments(value?: string) {
const documents = splitLines(value ?? '');
return documents.length ? documents : ['AI Gateway 提供 OpenAI 兼容接口。', '图片生成任务支持异步队列。'];
}
function splitLines(value: string) {
return value
.split(/\n+/)
.map((item) => item.trim())
.filter(Boolean);
}
function modelTypeForKind(kind: TaskForm['kind']) {
if (kind === 'embeddings') return 'text_embedding';
if (kind === 'reranks') return 'text_rerank';
if (kind === 'images.generations') return 'image_generate';
if (kind === 'images.edits') return 'image_edit';
if (kind === 'videos.generations') return 'video_generate';
if (kind === 'tasks.retrieve') return 'task';
return 'text_generate';
}