fix lock logic

This commit is contained in:
被遗忘的记忆
2024-12-12 00:57:11 +08:00
committed by GitHub
parent f8e576adb7
commit 5e24a37af1
+196 -55
View File
@@ -112,7 +112,14 @@ app.post("/v1/chat/completions", OpenAIApiKeyAuth, (req, res) => {
console.log("处理 OpenAI 格式的请求"); console.log("处理 OpenAI 格式的请求");
res.setHeader("Content-Type", "text/event-stream;charset=utf-8"); res.setHeader("Content-Type", "text/event-stream;charset=utf-8");
res.setHeader("Access-Control-Allow-Origin", "*"); res.setHeader("Access-Control-Allow-Origin", "*");
let jsonBody = JSON.parse(req.rawBody);
let jsonBody;
try {
jsonBody = JSON.parse(req.rawBody);
} catch (error) {
res.status(400).json({ error: { code: 400, message: "Invalid JSON" } });
return;
}
// 规范化消息 // 规范化消息
jsonBody.messages = await openaiNormalizeMessages(jsonBody.messages); jsonBody.messages = await openaiNormalizeMessages(jsonBody.messages);
@@ -129,22 +136,46 @@ app.post("/v1/chat/completions", OpenAIApiKeyAuth, (req, res) => {
} }
console.log("Using model " + jsonBody.model); console.log("Using model " + jsonBody.model);
// 调用 provider 获取回复 let selectedSession;
let releaseSessionCalled = false;
let completion;
let cancel;
// 定义释放会话
const releaseSession = () => {
if (selectedSession && !releaseSessionCalled) {
sessionManager.releaseSession(selectedSession);
console.log(`释放会话 ${selectedSession}`);
releaseSessionCalled = true;
}
};
// 监听客户端关闭事件
res.on("close", () => {
console.log(" > [Client closed]");
clientState.setClosed(true);
if (completion) {
completion.removeAllListeners();
}
if (cancel) {
cancel();
}
releaseSession();
});
try { try {
// 检查是否有可用会话 // 获取并锁定可用会话
const selectedSession = sessionManager.getSessionByStrategy('random'); const { selectedUsername, modeSwitched } = await sessionManager.getSessionByStrategy('round_robin');
selectedSession = selectedUsername;
console.log("Using session " + selectedSession); console.log("Using session " + selectedSession);
const {completion, cancel} = await provider.getCompletion({
({ completion, cancel } = await provider.getCompletion({
username: selectedSession, username: selectedSession,
messages: jsonBody.messages, messages: jsonBody.messages,
stream: !!jsonBody.stream, stream: !!jsonBody.stream,
proxyModel: jsonBody.model, proxyModel: jsonBody.model,
useCustomMode: process.env.USE_CUSTOM_MODE === "true" useCustomMode: process.env.USE_CUSTOM_MODE === "true",
}); modeSwitched: modeSwitched // 传递模式切换标志
// 释放账号 }));
const releaseSession = () => {
sessionManager.releaseSession(selectedSession);
};
// 监听开始事件 // 监听开始事件
completion.on("start", (id) => { completion.on("start", (id) => {
@@ -223,6 +254,7 @@ app.post("/v1/chat/completions", OpenAIApiKeyAuth, (req, res) => {
}) })
); );
res.end(); res.end();
releaseSession();
} }
}); });
@@ -235,22 +267,11 @@ app.post("/v1/chat/completions", OpenAIApiKeyAuth, (req, res) => {
releaseSession(); releaseSession();
}); });
// 监听客户端关闭事件
res.on("close", () => {
console.log(" > [Client closed]");
clientState.setClosed(true);
completion.removeAllListeners();
cancel();
releaseSession();
});
// 监听错误事件 // 监听错误事件
completion.on("error", releaseSession); completion.on("error", (err) => {
console.error("Completion error:", err);
} catch (error) { const errorMessage = "Error occurred: " + (err.message || "Unknown error");
console.error(error); if (!res.headersSent) {
const errorMessage = "Error occurred, please check the log.\n\n出现错误,请检查日志:<pre>" + (error.stack || error) + "</pre>";
res.status(500).send("No available sessions. 请检查你的账号状态。");
if (jsonBody.stream) { if (jsonBody.stream) {
res.write( res.write(
createEvent("data", { createEvent("data", {
@@ -274,6 +295,8 @@ app.post("/v1/chat/completions", OpenAIApiKeyAuth, (req, res) => {
system_fingerprint: "114514", system_fingerprint: "114514",
}) })
); );
res.write(createEvent("data", "[DONE]"));
res.end();
} else { } else {
res.write( res.write(
JSON.stringify({ JSON.stringify({
@@ -300,9 +323,72 @@ app.post("/v1/chat/completions", OpenAIApiKeyAuth, (req, res) => {
}, },
}) })
); );
}
res.end(); res.end();
}
}
releaseSession();
});
} catch (error) {
console.error("Request error:", error);
releaseSession();
const errorMessage = "Error occurred, please check the log.\n\n出现错误,请检查日志:<pre>" + (error.stack || error) + "</pre>";
if (!res.headersSent) {
if (jsonBody.stream) {
res.write(
createEvent("data", {
choices: [
{
content_filter_results: {
hate: { filtered: false, severity: "safe" },
self_harm: { filtered: false, severity: "safe" },
sexual: { filtered: false, severity: "safe" },
violence: { filtered: false, severity: "safe" },
},
delta: { content: errorMessage },
finish_reason: null,
index: 0,
},
],
created: Math.floor(new Date().getTime() / 1000),
id: uuidv4(),
model: jsonBody.model,
object: "chat.completion.chunk",
system_fingerprint: "114514",
})
);
res.write(createEvent("data", "[DONE]"));
res.end();
} else {
res.write(
JSON.stringify({
id: uuidv4(),
object: "chat.completion",
created: Math.floor(new Date().getTime() / 1000),
model: jsonBody.model,
system_fingerprint: "114514",
choices: [
{
index: 0,
message: {
role: "assistant",
content: errorMessage,
},
logprobs: null,
finish_reason: "stop",
},
],
usage: {
prompt_tokens: 1,
completion_tokens: 1,
total_tokens: 1,
},
})
);
res.end();
}
}
} }
}); });
}); });
@@ -420,7 +506,14 @@ app.post("/v1/messages", AnthropicApiKeyAuth, (req, res) => {
console.log("处理 Anthropic 格式的请求"); console.log("处理 Anthropic 格式的请求");
res.setHeader("Content-Type", "text/event-stream;charset=utf-8"); res.setHeader("Content-Type", "text/event-stream;charset=utf-8");
res.setHeader("Access-Control-Allow-Origin", "*"); res.setHeader("Access-Control-Allow-Origin", "*");
let jsonBody = JSON.parse(req.rawBody); let jsonBody;
try {
jsonBody = JSON.parse(req.rawBody);
} catch (error) {
res.status(400).json({ error: { code: 400, message: "Invalid JSON" } });
return;
}
// 处理消息格式 // 处理消息格式
jsonBody.messages = anthropicNormalizeMessages(jsonBody.messages); jsonBody.messages = anthropicNormalizeMessages(jsonBody.messages);
@@ -442,24 +535,49 @@ app.post("/v1/messages", AnthropicApiKeyAuth, (req, res) => {
} }
console.log(`Using model ${proxyModel}`); console.log(`Using model ${proxyModel}`);
// call provider to get completion let selectedSession;
let releaseSessionCalled = false;
let completion;
let cancel;
// 定义释放会话
const releaseSession = () => {
if (selectedSession && !releaseSessionCalled) {
sessionManager.releaseSession(selectedSession);
console.log(`释放会话 ${selectedSession}`);
releaseSessionCalled = true;
}
};
// 监听客户端关闭事件
res.on("close", () => {
console.log(" > [Client closed]");
clientState.setClosed(true);
if (completion) {
completion.removeAllListeners();
}
if (cancel) {
cancel();
}
releaseSession();
});
try { try {
// 检查是否有可用的会话 // 获取并锁定会话
const selectedSession = sessionManager.getSessionByStrategy('random'); const { selectedUsername, modeSwitched } = await sessionManager.getSessionByStrategy('round_robin');
selectedSession = selectedUsername;
console.log("Using session " + selectedSession); console.log("Using session " + selectedSession);
const {completion, cancel} = await provider.getCompletion({
({ completion, cancel } = await provider.getCompletion({
username: selectedSession, username: selectedSession,
messages: jsonBody.messages, messages: jsonBody.messages,
stream: !!jsonBody.stream, stream: !!jsonBody.stream,
proxyModel: proxyModel, proxyModel: proxyModel,
useCustomMode: process.env.USE_CUSTOM_MODE === "true" useCustomMode: process.env.USE_CUSTOM_MODE === "true",
}); modeSwitched: modeSwitched // 传递模式切换标志
}));
// 释放账号
const releaseSession = () => {
sessionManager.releaseSession(selectedSession);
};
// 监听开始事件
completion.on("start", (id) => { completion.on("start", (id) => {
if (jsonBody.stream) { if (jsonBody.stream) {
// send message start // send message start
@@ -485,6 +603,7 @@ app.post("/v1/messages", AnthropicApiKeyAuth, (req, res) => {
} }
}); });
// 监听完成事件
completion.on("completion", (id, text) => { completion.on("completion", (id, text) => {
if (jsonBody.stream) { if (jsonBody.stream) {
// send message delta // send message delta
@@ -507,9 +626,11 @@ app.post("/v1/messages", AnthropicApiKeyAuth, (req, res) => {
usage: { input_tokens: 0, output_tokens: 0 }, usage: { input_tokens: 0, output_tokens: 0 },
})); }));
res.end(); res.end();
releaseSession();
} }
}); });
// 监听结束事件
completion.on("end", () => { completion.on("end", () => {
if (jsonBody.stream) { if (jsonBody.stream) {
res.write(createEvent("content_block_stop", { type: "content_block_stop", index: 0 })); res.write(createEvent("content_block_stop", { type: "content_block_stop", index: 0 }));
@@ -524,28 +645,19 @@ app.post("/v1/messages", AnthropicApiKeyAuth, (req, res) => {
releaseSession(); releaseSession();
}); });
// 监听客户端关闭事件
res.on("close", () => {
console.log(" > [Client closed]");
clientState.setClosed(true);
completion.removeAllListeners();
cancel();
releaseSession();
});
// 监听错误事件 // 监听错误事件
completion.on("error", releaseSession); completion.on("error", (err) => {
console.error("Completion error:", err);
} catch (error) { // 向客户端返回错误信息
console.error(error); const errorMessage = "Error occurred: " + (err.message || "Unknown error");
const errorMessage = "Error occurred, please check the log.\\n\\n出现错误,请检查日志:<pre>" + (error.stack || error) + "</pre>"; if (!res.headersSent) {
res.status(500).send("No available sessions. 请检查你的账号状态。");
if (jsonBody.stream) { if (jsonBody.stream) {
res.write(createEvent("content_block_delta", { res.write(createEvent("content_block_delta", {
type: "content_block_delta", type: "content_block_delta",
index: 0, index: 0,
delta: { type: "text_delta", text: errorMessage }, delta: { type: "text_delta", text: errorMessage },
})); }));
res.end();
} else { } else {
res.write(JSON.stringify({ res.write(JSON.stringify({
id: uuidv4(), id: uuidv4(),
@@ -555,9 +667,38 @@ app.post("/v1/messages", AnthropicApiKeyAuth, (req, res) => {
stop_sequence: null, stop_sequence: null,
usage: { input_tokens: 0, output_tokens: 0 }, usage: { input_tokens: 0, output_tokens: 0 },
})); }));
}
res.end(); res.end();
} }
}
releaseSession();
});
} catch (error) {
console.error("Request error:", error);
releaseSession();
const errorMessage = "Error occurred, please check the log.\n\n出现错误,请检查日志:<pre>" + (error.stack || error) + "</pre>";
if (!res.headersSent) {
if (jsonBody.stream) {
res.write(createEvent("content_block_delta", {
type: "content_block_delta",
index: 0,
delta: { type: "text_delta", text: errorMessage },
}));
res.end();
} else {
res.write(JSON.stringify({
id: uuidv4(),
content: [{ text: errorMessage }, { id: "string", name: "string", input: {} }],
model: proxyModel,
stop_reason: "error",
stop_sequence: null,
usage: { input_tokens: 0, output_tokens: 0 },
}));
res.end();
}
}
}
}); });
}); });