210 lines
6.6 KiB
TypeScript
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';
|
|
}
|