Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion dto/openai_request.go
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ type GeneralOpenAIRequest struct {
Input any `json:"input,omitempty"`
Instruction string `json:"instruction,omitempty"`
Size string `json:"size,omitempty"`
Seconds *int `json:"seconds,omitempty"`
Seconds *string `json:"seconds,omitempty"`
Quality *string `json:"quality,omitempty"`
Functions json.RawMessage `json:"functions,omitempty"`
FrequencyPenalty *float64 `json:"frequency_penalty,omitempty"`
Expand Down
1 change: 1 addition & 0 deletions web/src/constants/playground.constants.js
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,7 @@ export const DEBUG_TABS = {
// ========== API 相关常量 ==========
export const API_ENDPOINTS = {
CHAT_COMPLETIONS: '/pg/chat/completions',
VIDEO_GENERATIONS: '/v1/video/generations',
USER_MODELS: '/api/user/models',
USER_GROUPS: '/api/user/self/groups',
};
Expand Down
6 changes: 2 additions & 4 deletions web/src/helpers/api.js
Original file line number Diff line number Diff line change
Expand Up @@ -158,14 +158,12 @@ export const buildApiPayload = (
const isVideoModel =
typeof inputs.model === 'string' && inputs.model.includes('video');
if (isVideoModel) {
payload.stream = false;
if (inputs.videoSize) {
payload.size = inputs.videoSize;
}
if (inputs.videoSeconds) {
const parsedSeconds = Number.parseInt(inputs.videoSeconds, 10);
if (Number.isFinite(parsedSeconds)) {
payload.seconds = parsedSeconds;
}
payload.seconds = String(inputs.videoSeconds);
}
if (inputs.videoQuality) {
payload.quality = inputs.videoQuality;
Expand Down
125 changes: 119 additions & 6 deletions web/src/hooks/playground/useApiRequest.jsx
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,78 @@ export const useApiRequest = (
saveMessages,
) => {
const { t } = useTranslation();
const isVideoGenerationPayload = useCallback((payload) => {
const model = payload?.model;
return typeof model === 'string' && model.includes('video');
}, []);

const getTextFromMessageContent = useCallback((content) => {
if (typeof content === 'string') {
return content;
}
if (!Array.isArray(content)) {
return '';
}
const textParts = content
.filter((item) => item?.type === 'text')
.map((item) => item?.text || '')
.filter(Boolean);
return textParts.join('\n');
}, []);

const getImageFromMessageContent = useCallback((content) => {
if (!Array.isArray(content)) {
return '';
}
const imageItem = content.find((item) => item?.type === 'image_url');
if (!imageItem) {
return '';
}
const imageURL = imageItem.image_url;
if (typeof imageURL === 'string') {
return imageURL;
}
return imageURL?.url || '';
}, []);

const buildVideoRequestPayload = useCallback(
(payload) => {
const messages = Array.isArray(payload?.messages) ? payload.messages : [];
const lastUserMessage = [...messages]
.reverse()
.find((m) => m?.role === 'user');
const prompt = getTextFromMessageContent(lastUserMessage?.content);
const image = getImageFromMessageContent(lastUserMessage?.content);

return {
model: payload.model,
prompt,
seconds: payload.seconds,
size: payload.size,
quality: payload.quality,
...(image ? { image } : {}),
};
},
[getImageFromMessageContent, getTextFromMessageContent],
);

const resolveEndpointAndPayload = useCallback(
(payload) => {
if (isVideoGenerationPayload(payload)) {
return {
endpoint: API_ENDPOINTS.VIDEO_GENERATIONS,
requestPayload: buildVideoRequestPayload(payload),
forceNonStream: true,
};
}
return {
endpoint: API_ENDPOINTS.CHAT_COMPLETIONS,
requestPayload: payload,
forceNonStream: false,
};
},
[buildVideoRequestPayload, isVideoGenerationPayload],
);

// 处理消息自动关闭逻辑的公共函数
const applyAutoCollapseLogic = useCallback(
Expand Down Expand Up @@ -174,9 +246,10 @@ export const useApiRequest = (
// 非流式请求
const handleNonStreamRequest = useCallback(
async (payload) => {
const { endpoint, requestPayload } = resolveEndpointAndPayload(payload);
setDebugData((prev) => ({
...prev,
request: payload,
request: requestPayload,
timestamp: new Date().toISOString(),
response: null,
sseMessages: null, // 非流式请求清除 SSE 消息
Expand All @@ -185,13 +258,13 @@ export const useApiRequest = (
setActiveDebugTab(DEBUG_TABS.REQUEST);

try {
const response = await fetch(API_ENDPOINTS.CHAT_COMPLETIONS, {
const response = await fetch(endpoint, {
method: 'POST',
headers: {
'Content-Type': 'application/json',
'New-Api-User': getUserIdFromLocalStorage(),
},
body: JSON.stringify(payload),
body: JSON.stringify(requestPayload),
});

if (!response.ok) {
Expand Down Expand Up @@ -228,6 +301,38 @@ export const useApiRequest = (
}));
setActiveDebugTab(DEBUG_TABS.RESPONSE);

if (
endpoint === API_ENDPOINTS.VIDEO_GENERATIONS ||
data.object === 'video' ||
data.task_id
) {
const summary = [
`${t('视频任务已创建')}`,
`task_id: ${data.task_id || data.id || '-'}`,
`status: ${data.status || '-'}`,
`seconds: ${data.seconds || requestPayload.seconds || '-'}`,
`size: ${data.size || requestPayload.size || '-'}`,
].join('\n');
setMessage((prevMessage) => {
const newMessages = [...prevMessage];
const lastMessage = newMessages[newMessages.length - 1];
if (lastMessage?.status === MESSAGE_STATUS.LOADING) {
const autoCollapseState = applyAutoCollapseLogic(
lastMessage,
true,
);
newMessages[newMessages.length - 1] = {
...lastMessage,
content: summary,
status: MESSAGE_STATUS.COMPLETE,
...autoCollapseState,
};
}
return newMessages;
});
return;
}

if (data.choices?.[0]) {
const choice = data.choices[0];
let content = choice.message?.content || '';
Expand Down Expand Up @@ -285,7 +390,14 @@ export const useApiRequest = (
});
}
},
[setDebugData, setActiveDebugTab, setMessage, t, applyAutoCollapseLogic],
[
resolveEndpointAndPayload,
setDebugData,
setActiveDebugTab,
setMessage,
t,
applyAutoCollapseLogic,
],
);

// SSE请求
Expand Down Expand Up @@ -500,13 +612,14 @@ export const useApiRequest = (
// 发送请求
const sendRequest = useCallback(
(payload, isStream) => {
if (isStream) {
const { forceNonStream } = resolveEndpointAndPayload(payload);
if (isStream && !forceNonStream) {
handleSSE(payload);
} else {
handleNonStreamRequest(payload);
}
},
[handleSSE, handleNonStreamRequest],
[resolveEndpointAndPayload, handleSSE, handleNonStreamRequest],
);

return {
Expand Down