225 lines
7.0 KiB
TypeScript
225 lines
7.0 KiB
TypeScript
// server/api/search.ts - OCS AnswererWrapper 兼容搜索接口
|
||
import {
|
||
getQuery,
|
||
type H3Event,
|
||
readBody,
|
||
readFormData,
|
||
setResponseStatus
|
||
} from "h3";
|
||
|
||
import {
|
||
ANSWER_SYSTEM_PROMPT,
|
||
createOcsErrorResponse,
|
||
extractAnswer,
|
||
formatAnswerForOcs,
|
||
parseQuestionAndOptions,
|
||
type SearchParams
|
||
} from "~~/server/utils/answer";
|
||
import { getAuthSession } from "~~/server/utils/auth";
|
||
import { answerCache } from "~~/server/utils/cache";
|
||
import { prisma } from "~~/server/utils/db";
|
||
import { createApiLogger, toSafeLogError } from "~~/server/utils/logging";
|
||
import { askAnswerStream } from "~~/server/utils/openai";
|
||
import { addQaRecord } from "~~/server/utils/runtimeState";
|
||
|
||
/**
|
||
* 把 query/body/form 中的值统一转成字符串
|
||
*
|
||
* OCS 传参通常是普通字符串,但 query 可能出现同名参数数组;
|
||
* 为了兼容旧 Flask 服务,这里取第一个值并把缺失值转为空字符串
|
||
*/
|
||
const toStringValue = (value: unknown) => {
|
||
if (Array.isArray(value)) return value[0]?.toString() || "";
|
||
return value?.toString() || "";
|
||
};
|
||
|
||
/** 判断 JSON body 是否是普通对象;数组和字符串都不是本接口接受的 body */
|
||
const isRecord = (value: unknown): value is Record<string, unknown> => {
|
||
return typeof value === "object" && value !== null && !Array.isArray(value);
|
||
};
|
||
|
||
/** SearchParams 扩展,包含 OCS 脚本通过 body 传入的 apiToken */
|
||
type SearchParamsWithToken = SearchParams & { token?: string };
|
||
|
||
/**
|
||
* 兼容旧 Python 服务的三种入参方式
|
||
*
|
||
* 1. GET:从 query 读取 `title/type/options`
|
||
* 2. multipart/form-data:从 FormData 读取
|
||
* 3. JSON 或 urlencoded POST:通过 `readBody` 读取对象
|
||
*
|
||
* 返回 `"invalid_body"` 时,说明客户端传了数组、字符串等不合法 body,
|
||
* handler 会按安全规范返回 400 和本地错误文案
|
||
*/
|
||
const readSearchParams = async (
|
||
event: H3Event
|
||
): Promise<SearchParamsWithToken | "invalid_body"> => {
|
||
const method = event.node.req.method?.toUpperCase() || "GET";
|
||
|
||
if (method === "GET") {
|
||
const query = getQuery(event);
|
||
|
||
return {
|
||
title: toStringValue(query.title).trim(),
|
||
type: toStringValue(query.type).trim(),
|
||
options: toStringValue(query.options).trim(),
|
||
token: toStringValue(query.token).trim() || undefined
|
||
};
|
||
}
|
||
|
||
const contentType = event.node.req.headers["content-type"] || "";
|
||
|
||
if (contentType.includes("multipart/form-data")) {
|
||
const form = await readFormData(event);
|
||
|
||
return {
|
||
title: toStringValue(form.get("title")).trim(),
|
||
type: toStringValue(form.get("type")).trim(),
|
||
options: toStringValue(form.get("options")).trim()
|
||
};
|
||
}
|
||
|
||
const body = await readBody<unknown>(event).catch(() => null);
|
||
if (body === null || body === undefined) {
|
||
return {
|
||
title: "",
|
||
type: "",
|
||
options: ""
|
||
};
|
||
}
|
||
|
||
if (!isRecord(body)) {
|
||
return "invalid_body";
|
||
}
|
||
|
||
return {
|
||
title: toStringValue(body.title).trim(),
|
||
type: toStringValue(body.type).trim(),
|
||
options: toStringValue(body.options).trim(),
|
||
token: toStringValue(body.token).trim() || undefined
|
||
};
|
||
};
|
||
|
||
/**
|
||
* OCS 答题搜索主接口
|
||
*
|
||
* 流程:
|
||
* - 校验请求方法
|
||
* - 鉴权:优先读 session(浏览器登录),无 session 则读 body.token(OCS 油猴脚本)
|
||
* - 两种识别方式均无效时返回 401
|
||
* - 读取题目、题型、选项,兼容 GET/JSON/form
|
||
* - 先查内存缓存,命中后不再请求 OpenAI
|
||
* - 未命中时拼提示词,服务端流式请求 Chat Completions
|
||
* - 清洗答案、写入缓存和 DB 记录,最后返回 OCS 兼容 JSON
|
||
*
|
||
* 安全边界:
|
||
* - 不向前端返回 OpenAI 错误、响应体、API Key 或内部堆栈
|
||
* - 日志只记录长度、阶段、耗时等摘要,不记录完整题目和完整 prompt
|
||
* - apiToken 不进日志
|
||
*/
|
||
export default defineEventHandler(async (event) => {
|
||
const logger = createApiLogger(event, "api.search");
|
||
const method = event.node.req.method?.toUpperCase() || "GET";
|
||
const startedAt = Date.now();
|
||
|
||
// Nitro 文件路由会匹配所有方法;这里显式限制,保持接口行为清楚
|
||
if (!["GET", "POST"].includes(method)) {
|
||
setResponseStatus(event, 405);
|
||
return createOcsErrorResponse("请求方法不支持");
|
||
}
|
||
|
||
try {
|
||
// 读取并标准化 OCS 参数,避免后续逻辑关心请求来源
|
||
const params = await readSearchParams(event);
|
||
if (params === "invalid_body") {
|
||
setResponseStatus(event, 400);
|
||
return createOcsErrorResponse("请求体必须是 JSON 对象");
|
||
}
|
||
|
||
// 鉴权:优先 session(浏览器登录),无 session 则用 body.token 查 apiToken
|
||
// OCS 油猴脚本跨域无法携带 cookie,需在 body 中传入用户自己的 apiToken
|
||
let userId: string | null = null;
|
||
|
||
const session = await getAuthSession(event);
|
||
if (session) {
|
||
userId = session.user.id;
|
||
} else if (params.token) {
|
||
const user = await prisma.user.findUnique({
|
||
where: { apiToken: params.token },
|
||
select: { id: true }
|
||
});
|
||
userId = user?.id ?? null;
|
||
}
|
||
|
||
if (!userId) {
|
||
setResponseStatus(event, 401);
|
||
logger.warn("unauthorized");
|
||
return createOcsErrorResponse("请先登录或提供有效 token");
|
||
}
|
||
|
||
logger.info("read_question", {
|
||
questionLength: params.title.length,
|
||
type: params.type,
|
||
hasOptions: Boolean(params.options)
|
||
});
|
||
|
||
if (!params.title) {
|
||
return createOcsErrorResponse("未提供问题内容");
|
||
}
|
||
|
||
// 缓存 key 包含题目、题型和选项;同题不同选项不能共用答案
|
||
const cachedAnswer = answerCache?.get(
|
||
params.title,
|
||
params.type,
|
||
params.options
|
||
);
|
||
|
||
if (cachedAnswer) {
|
||
logger.info("cache_hit");
|
||
return formatAnswerForOcs(params.title, cachedAnswer);
|
||
}
|
||
|
||
// 构造和旧 Python 服务一致的提示词,再通过服务端流式请求模型
|
||
const prompt = parseQuestionAndOptions(
|
||
params.title,
|
||
params.options,
|
||
params.type
|
||
);
|
||
const streamResult = await askAnswerStream({
|
||
prompt,
|
||
systemPrompt: ANSWER_SYSTEM_PROMPT
|
||
});
|
||
const processedAnswer = extractAnswer(streamResult.answer, params.type);
|
||
|
||
// 先缓存再记录;这两步失败风险很低,且都是内存操作,不会阻塞主链路
|
||
answerCache?.set(
|
||
params.title,
|
||
processedAnswer,
|
||
params.type,
|
||
params.options
|
||
);
|
||
await addQaRecord({
|
||
userId,
|
||
question: params.title,
|
||
type: params.type,
|
||
options: params.options,
|
||
answer: processedAnswer
|
||
});
|
||
|
||
logger.info("finish_success", {
|
||
durationMs: Date.now() - startedAt,
|
||
chunkCount: streamResult.upstreamResponse.chunkCount,
|
||
answerLength: processedAnswer.length
|
||
});
|
||
|
||
return formatAnswerForOcs(params.title, processedAnswer);
|
||
} catch (error) {
|
||
// 按安全优先策略:详细错误只进服务端日志,OCS 端只看到本地通用文案
|
||
logger.error("finish_error", {
|
||
error: toSafeLogError(error)
|
||
});
|
||
|
||
return createOcsErrorResponse("服务器内部错误");
|
||
}
|
||
});
|