|
@@ -0,0 +1,139 @@
|
|
|
|
|
+package com.zsjz.ai.module.agent.service;
|
|
|
|
|
+
|
|
|
|
|
+import com.zsjz.ai.common.enums.ModelTypeEnum;
|
|
|
|
|
+import com.zsjz.ai.common.exception.ServerException;
|
|
|
|
|
+import com.zsjz.ai.module.agent.entity.AgentModel;
|
|
|
|
|
+import com.zsjz.ai.module.agent.rag.EmbeddingModelFactory;
|
|
|
|
|
+import com.zsjz.ai.module.agent.vo.ModelTestResultVO;
|
|
|
|
|
+import io.agentscope.core.message.Msg;
|
|
|
|
|
+import io.agentscope.core.message.MsgRole;
|
|
|
|
|
+import io.agentscope.core.message.TextBlock;
|
|
|
|
|
+import io.agentscope.core.model.GenerateOptions;
|
|
|
|
|
+import io.agentscope.core.model.Model;
|
|
|
|
|
+import lombok.RequiredArgsConstructor;
|
|
|
|
|
+import lombok.extern.slf4j.Slf4j;
|
|
|
|
|
+import org.springframework.stereotype.Component;
|
|
|
|
|
+import org.springframework.util.StringUtils;
|
|
|
|
|
+
|
|
|
|
|
+import java.time.Duration;
|
|
|
|
|
+import java.util.List;
|
|
|
|
|
+import java.util.concurrent.TimeoutException;
|
|
|
|
|
+
|
|
|
|
|
+/**
|
|
|
|
|
+ * 模型连通性测试:对已保存的模型配置发起一次<b>最小化真实调用</b>,
|
|
|
|
|
+ * 验证 Base URL / API Key / 模型 ID 组合是否可用。
|
|
|
|
|
+ *
|
|
|
|
|
+ * <p>对话模型取流式响应的首个分片即算畅通(不等完整回复,省 token 也更快);
|
|
|
|
|
+ * 向量模型对固定短文本做一次向量化。工厂抛出的配置校验异常(缺模型 ID / 缺 API Key)
|
|
|
|
|
+ * 与真实业务链路(Agent / RAG)同源,原样带出原因文案。
|
|
|
|
|
+ */
|
|
|
|
|
+@Slf4j
|
|
|
|
|
+@Component
|
|
|
|
|
+@RequiredArgsConstructor
|
|
|
|
|
+public class ModelConnectivityTester {
|
|
|
|
|
+
|
|
|
|
|
+ /**
|
|
|
|
|
+ * 单次测试超时:比业务调用短 —— 连通性测试要快速给出结论;
|
|
|
|
|
+ * 又不能太短,本地 Ollama 冷启动加载模型可能要十几秒
|
|
|
|
|
+ */
|
|
|
|
|
+ private static final Duration TEST_TIMEOUT = Duration.ofSeconds(20);
|
|
|
|
|
+
|
|
|
|
|
+ /**
|
|
|
|
|
+ * 测试调用的输出上限,防止一次 ping 生成一大段回复
|
|
|
|
|
+ */
|
|
|
|
|
+ private static final int PING_MAX_TOKENS = 16;
|
|
|
|
|
+
|
|
|
|
|
+ /**
|
|
|
|
|
+ * 失败信息最大长度(底层异常链的消息可能很长,前端只展示摘要)
|
|
|
|
|
+ */
|
|
|
|
|
+ private static final int MAX_MESSAGE_LEN = 200;
|
|
|
|
|
+
|
|
|
|
|
+ private static final String PING_TEXT = "ping";
|
|
|
|
|
+
|
|
|
|
|
+ private final AgentModelFactory agentModelFactory;
|
|
|
|
|
+ private final EmbeddingModelFactory embeddingModelFactory;
|
|
|
|
|
+
|
|
|
|
|
+ /**
|
|
|
|
|
+ * 对给定模型配置执行连通性测试。
|
|
|
|
|
+ *
|
|
|
|
|
+ * <p>连通失败是本接口的<b>预期结果</b>而非服务端故障,
|
|
|
|
|
+ * 因此失败详情放在返回值里(success=false + message),不抛异常,
|
|
|
|
|
+ * 前端拿到的是 HTTP 200 的结构化结论而非报错弹窗。
|
|
|
|
|
+ *
|
|
|
|
|
+ * @param config 模型配置(model 表记录)
|
|
|
|
|
+ * @return 测试结果(是否畅通 + 耗时 + 失败原因)
|
|
|
|
|
+ */
|
|
|
|
|
+ public ModelTestResultVO test(AgentModel config) {
|
|
|
|
|
+ long start = System.currentTimeMillis();
|
|
|
|
|
+ try {
|
|
|
|
|
+ if (ModelTypeEnum.EMBEDDING.getCode().equals(config.getType())) {
|
|
|
|
|
+ testEmbedding(config);
|
|
|
|
|
+ } else {
|
|
|
|
|
+ testChat(config);
|
|
|
|
|
+ }
|
|
|
|
|
+ long latency = System.currentTimeMillis() - start;
|
|
|
|
|
+ log.info("模型连通性测试成功: name={}, 耗时={}ms", config.getName(), latency);
|
|
|
|
|
+ ModelTestResultVO vo = new ModelTestResultVO();
|
|
|
|
|
+ vo.setSuccess(true);
|
|
|
|
|
+ vo.setLatencyMs(latency);
|
|
|
|
|
+ return vo;
|
|
|
|
|
+ } catch (ServerException e) {
|
|
|
|
|
+ return failure(config, start, e.getMessage());
|
|
|
|
|
+ } catch (Exception e) {
|
|
|
|
|
+ return failure(config, start, describeFailure(e));
|
|
|
|
|
+ }
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ /**
|
|
|
|
|
+ * 对话模型测试:发一条 "ping",拿到首个流式分片即取消
|
|
|
|
|
+ */
|
|
|
|
|
+ private void testChat(AgentModel config) {
|
|
|
|
|
+ Model model = agentModelFactory.create(config);
|
|
|
|
|
+ List<Msg> messages = List.of(Msg.builder().role(MsgRole.USER).textContent(PING_TEXT).build());
|
|
|
|
|
+ GenerateOptions options = GenerateOptions.builder().maxTokens(PING_MAX_TOKENS).build();
|
|
|
|
|
+ model.stream(messages, List.of(), options)
|
|
|
|
|
+ .take(1)
|
|
|
|
|
+ .timeout(TEST_TIMEOUT)
|
|
|
|
|
+ .blockLast();
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ /**
|
|
|
|
|
+ * 向量模型测试:对固定短文本做一次向量化
|
|
|
|
|
+ */
|
|
|
|
|
+ private void testEmbedding(AgentModel config) {
|
|
|
|
|
+ EmbeddingModelFactory.EmbeddingSpec spec = embeddingModelFactory.create(config);
|
|
|
|
|
+ double[] vec = spec.model().embed(TextBlock.builder().text(PING_TEXT).build())
|
|
|
|
|
+ .block(TEST_TIMEOUT);
|
|
|
|
|
+ if (vec == null || vec.length == 0) {
|
|
|
|
|
+ throw new ServerException(500, "嵌入服务返回空向量");
|
|
|
|
|
+ }
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ private ModelTestResultVO failure(AgentModel config, long start, String message) {
|
|
|
|
|
+ long latency = System.currentTimeMillis() - start;
|
|
|
|
|
+ log.warn("模型连通性测试失败: name={}, 耗时={}ms, 原因={}", config.getName(), latency, message);
|
|
|
|
|
+ ModelTestResultVO vo = new ModelTestResultVO();
|
|
|
|
|
+ vo.setSuccess(false);
|
|
|
|
|
+ vo.setLatencyMs(latency);
|
|
|
|
|
+ vo.setMessage(StringUtils.hasText(message) ? message : "连接失败");
|
|
|
|
|
+ return vo;
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ /**
|
|
|
|
|
+ * 把底层异常翻译成用户可读的原因:超时单列,其余取根因消息
|
|
|
|
|
+ */
|
|
|
|
|
+ private static String describeFailure(Throwable e) {
|
|
|
|
|
+ for (Throwable t = e; t != null; t = (t.getCause() == t ? null : t.getCause())) {
|
|
|
|
|
+ if (t instanceof TimeoutException) {
|
|
|
|
|
+ return "连接超时(" + TEST_TIMEOUT.toSeconds() + "s),请检查网络与 Base URL";
|
|
|
|
|
+ }
|
|
|
|
|
+ }
|
|
|
|
|
+ Throwable root = e;
|
|
|
|
|
+ while (root.getCause() != null && root.getCause() != root) {
|
|
|
|
|
+ root = root.getCause();
|
|
|
|
|
+ }
|
|
|
|
|
+ String msg = StringUtils.hasText(root.getMessage()) ? root.getMessage()
|
|
|
|
|
+ : (StringUtils.hasText(e.getMessage()) ? e.getMessage() : e.getClass().getSimpleName());
|
|
|
|
|
+ return msg.length() > MAX_MESSAGE_LEN ? msg.substring(0, MAX_MESSAGE_LEN) + "..." : msg;
|
|
|
|
|
+ }
|
|
|
|
|
+}
|