Skip to content
Draft
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
12 changes: 12 additions & 0 deletions AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,18 @@ npm run build
npm run dev # 开发服务器 5173,/api 代理到 localhost:9090
```

### 最小相关检查路由

局部改动先跑最小相关检查获得快速反馈,提交前仍按影响范围升级到完整门禁。完整映射与升级条件见 `docs/verification-routing.md`。

```bash
# 后端单个测试类;-am 场景必须关闭“未找到指定测试即失败”
./mvnw -B -ntp -pl bootstrap -am -Dtest=StreamChatTraceRunnerTest -Dsurefire.failIfNoSpecifiedTests=false test

# 前端单个测试文件
cd frontend && npm run test -- src/hooks/__tests__/useStreamResponse.test.ts
```

## 模块分层

```text
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,5 +17,5 @@

package com.nageoffer.ai.ragent.rag.dto;

public record MetaPayload(String conversationId, String taskId) {
public record MetaPayload(String conversationId, String taskId, String traceId) {
}
Original file line number Diff line number Diff line change
Expand Up @@ -79,18 +79,24 @@ public StreamChatEventHandler(StreamChatHandlerParams params) {
this.messageChunkSize = resolveMessageChunkSize(params.getModelProperties());
this.sendTitleOnComplete = shouldSendTitle();

// 初始化(发送初始事件、注册任务)
// 先返回 taskId,保证排队期间也可取消;Trace 建立后会发送一次带 traceId 的 META 更新。
initialize();
}

/**
* 初始化:发送元数据事件并注册任务
* 初始化:发送可取消所需的元数据并注册任务
*/
private void initialize() {
sender.sendEvent(SSEEventType.META.value(), new MetaPayload(conversationId, taskId));
sender.sendEvent(SSEEventType.META.value(), new MetaPayload(conversationId, taskId, null));
taskManager.register(taskId, sender, this::buildCompletionPayloadOnCancel);
}

@Override
public void onTraceStarted(String traceId) {
taskManager.attachTrace(taskId, traceId);
sender.sendEvent(SSEEventType.META.value(), new MetaPayload(conversationId, taskId, traceId));
}

/**
* 解析消息块大小
*/
Expand Down Expand Up @@ -127,7 +133,9 @@ private CompletionPayload buildCompletionPayloadOnCancel() {
message.setMessageStatus(ChatMessage.MessageStatus.INTERRUPTED);
messageId = memoryService.append(conversationId, userId, message);
} catch (Exception e) {
log.error("取消时持久化消息失败,conversationId:{}", conversationId, e);
log.error("Failed to persist cancelled SSE message: traceId={}, taskId={}, errorType={}",
StreamTaskManager.safeCorrelationId(taskManager.traceId(taskId)),
StreamTaskManager.safeCorrelationId(taskId), e.getClass().getSimpleName());
}
}
String title = resolveTitleForEvent();
Expand Down Expand Up @@ -209,13 +217,18 @@ public void onComplete() {
message.setMessageStatus(ChatMessage.MessageStatus.NORMAL);
messageId = memoryService.append(conversationId, userId, message);
} catch (Exception e) {
log.error("对话完成时持久化消息失败,conversationId:{}", conversationId, e);
log.error("Failed to persist completed SSE message: traceId={}, taskId={}, errorType={}",
StreamTaskManager.safeCorrelationId(taskManager.traceId(taskId)),
StreamTaskManager.safeCorrelationId(taskId), e.getClass().getSimpleName());
}
String title = resolveTitleForEvent();
String messageIdText = StrUtil.isBlank(messageId) ? null : messageId;
sender.sendEvent(SSEEventType.FINISH.value(),
new CompletionPayload(messageIdText, title, sources, ChatMessage.MessageStatus.NORMAL));
sender.sendEvent(SSEEventType.DONE.value(), "[DONE]");
log.info("SSE stream completed: traceId={}, taskId={}",
StreamTaskManager.safeCorrelationId(taskManager.traceId(taskId)),
StreamTaskManager.safeCorrelationId(taskId));
taskManager.unregister(taskId);
sender.complete();
}
Expand All @@ -225,6 +238,9 @@ public void onError(Throwable t) {
if (taskManager.isCancelled(taskId)) {
return;
}
log.warn("SSE stream failed: traceId={}, taskId={}",
StreamTaskManager.safeCorrelationId(taskManager.traceId(taskId)),
StreamTaskManager.safeCorrelationId(taskId));
taskManager.unregister(taskId);
sender.fail(t);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,23 @@ public void register(String taskId, SseEmitterSender sender, Supplier<Completion
}
}

public void attachTrace(String taskId, String traceId) {
getOrCreate(taskId).traceId = traceId;
}

public String traceId(String taskId) {
StreamTaskInfo info = tasks.getIfPresent(taskId);
return info == null ? null : info.traceId;
}

public void bindCancellationObserver(String taskId, Runnable observer) {
StreamTaskInfo taskInfo = getOrCreate(taskId);
taskInfo.cancellationObserver = observer;
if (taskInfo.cancelled.get()) {
notifyCancellationObserver(taskInfo);
}
}

public void bindHandle(String taskId, StreamCancellationHandle handle) {
StreamTaskInfo taskInfo = getOrCreate(taskId);
taskInfo.handle = handle;
Expand All @@ -101,6 +118,7 @@ public boolean isCancelled(String taskId) {
}

public void cancel(String taskId) {
log.info("SSE cancellation requested: taskId={}", safeCorrelationId(taskId));
// 先设置 Redis 标记,再发布消息
RBucket<Boolean> bucket = redissonClient.getBucket(cancelKey(taskId));
bucket.set(Boolean.TRUE, CANCEL_TTL);
Expand Down Expand Up @@ -139,15 +157,21 @@ private void cancelLocal(String taskId) {
return;
}

if (taskInfo.handle != null) {
taskInfo.handle.cancel();
}
try {
if (taskInfo.handle != null) {
taskInfo.handle.cancel();
}

// 在取消时执行回调,保存已累积的内容
if (taskInfo.sender != null) {
CompletionPayload payload = taskInfo.onCancelSupplier.get();
sendCancelAndDone(taskInfo.sender, payload);
taskInfo.sender.complete();
// 在取消时执行回调,保存已累积的内容
if (taskInfo.sender != null) {
CompletionPayload payload = taskInfo.onCancelSupplier.get();
sendCancelAndDone(taskInfo.sender, payload);
taskInfo.sender.complete();
}
} finally {
notifyCancellationObserver(taskInfo);
log.info("SSE cancellation applied: traceId={}, taskId={}",
safeCorrelationId(taskInfo.traceId), safeCorrelationId(taskId));
}
}

Expand All @@ -169,15 +193,29 @@ private void sendCancelAndDone(SseEmitterSender sender, CompletionPayload payloa
sender.sendEvent(SSEEventType.DONE.value(), "[DONE]");
}

private void notifyCancellationObserver(StreamTaskInfo taskInfo) {
Runnable observer = taskInfo.cancellationObserver;
if (observer != null && taskInfo.cancellationObserved.compareAndSet(false, true)) {
observer.run();
}
}

public static String safeCorrelationId(String value) {
return value != null && value.matches("[A-Za-z0-9_-]{1,64}") ? value : "<unavailable>";
}

@SneakyThrows
private StreamTaskInfo getOrCreate(String taskId) {
return tasks.get(taskId, StreamTaskInfo::new);
}

private static final class StreamTaskInfo {
private final AtomicBoolean cancelled = new AtomicBoolean(false);
private final AtomicBoolean cancellationObserved = new AtomicBoolean(false);
private volatile StreamCancellationHandle handle;
private volatile SseEmitterSender sender;
private volatile Supplier<CompletionPayload> onCancelSupplier;
private volatile Runnable cancellationObserver;
private volatile String traceId;
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
import com.nageoffer.ai.ragent.framework.convention.ChatMessage;
import com.nageoffer.ai.ragent.framework.convention.ChatRequest;
import com.nageoffer.ai.ragent.framework.convention.SourceRef;
import com.nageoffer.ai.ragent.framework.trace.RagTraceContext;
import com.nageoffer.ai.ragent.infra.chat.LLMService;
import com.nageoffer.ai.ragent.infra.chat.StreamCallback;
import com.nageoffer.ai.ragent.infra.chat.StreamCancellationHandle;
Expand Down Expand Up @@ -80,23 +81,51 @@ public class StreamChatPipeline {
* 执行流式对话管道
*/
public void execute(StreamChatContext ctx) {
loadMemory(ctx);
rewriteQuery(ctx);
resolveIntents(ctx);
String stage = "memory";
try {
logStage(stage, ctx);
loadMemory(ctx);
stage = "rewrite";
logStage(stage, ctx);
rewriteQuery(ctx);
stage = "intent";
logStage(stage, ctx);
resolveIntents(ctx);

if (handleGuidance(ctx)) {
return;
}
if (handleSystemOnly(ctx)) {
return;
}
stage = "guidance";
logStage(stage, ctx);
if (handleGuidance(ctx)) {
return;
}
stage = "system-response";
logStage(stage, ctx);
if (handleSystemOnly(ctx)) {
return;
}

RetrievalContext retrievalCtx = retrieve(ctx);
if (handleEmptyRetrieval(ctx, retrievalCtx)) {
return;
stage = "retrieval";
logStage(stage, ctx);
RetrievalContext retrievalCtx = retrieve(ctx);
if (handleEmptyRetrieval(ctx, retrievalCtx)) {
return;
}

stage = "llm-stream";
logStage(stage, ctx);
streamRagResponse(ctx, retrievalCtx);
} catch (RuntimeException ex) {
log.warn("SSE pipeline failed: traceId={}, taskId={}, stage={}, errorType={}",
StreamTaskManager.safeCorrelationId(RagTraceContext.getTraceId()),
StreamTaskManager.safeCorrelationId(ctx.getTaskId()), stage,
ex.getClass().getSimpleName());
throw ex;
}
}

streamRagResponse(ctx, retrievalCtx);
private void logStage(String stage, StreamChatContext ctx) {
log.debug("SSE pipeline stage: traceId={}, taskId={}, stage={}",
StreamTaskManager.safeCorrelationId(RagTraceContext.getTraceId()),
StreamTaskManager.safeCorrelationId(ctx.getTaskId()), stage);
}

// ==================== 流水线阶段 ====================
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -145,7 +145,9 @@ private String buildFallbackTitle(String question) {
private void sendRejectEvents(SseEmitter emitter, RejectedContext rejectedContext) {
SseEmitterSender sender = new SseEmitterSender(emitter);
if (rejectedContext != null) {
sender.sendEvent(SSEEventType.META.value(), new MetaPayload(rejectedContext.conversationId, rejectedContext.taskId));
// 限流拒绝发生在 Trace 建立前,不能返回一个没有对应运行记录的伪 traceId。
sender.sendEvent(SSEEventType.META.value(),
new MetaPayload(rejectedContext.conversationId, rejectedContext.taskId, null));
sender.sendEvent(SSEEventType.REJECT.value(), new MessageDelta(RESPONSE_TYPE, REJECT_MESSAGE));
sender.sendEvent(SSEEventType.FINISH.value(),
new CompletionPayload(String.valueOf(rejectedContext.messageId), rejectedContext.title,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,11 +28,13 @@
import com.nageoffer.ai.ragent.rag.dao.entity.RagTraceNodeDO;
import com.nageoffer.ai.ragent.rag.dao.entity.RagTraceRunDO;
import com.nageoffer.ai.ragent.rag.service.RagTraceRecordService;
import com.nageoffer.ai.ragent.rag.service.handler.StreamTaskManager;
import lombok.RequiredArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import org.springframework.stereotype.Component;

import java.util.Date;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.function.Consumer;

/**
Expand All @@ -48,11 +50,13 @@ public class StreamChatTraceRunner {
private static final String STATUS_RUNNING = "RUNNING";
private static final String STATUS_SUCCESS = "SUCCESS";
private static final String STATUS_ERROR = "ERROR";
private static final String STATUS_CANCELLED = "CANCELLED";
private static final String USER_TTFT_NODE_NAME = "user-first-packet";
private static final String USER_TTFT_NODE_TYPE = "USER_TTFT";

private final RagTraceProperties traceProperties;
private final RagTraceRecordService traceRecordService;
private final StreamTaskManager taskManager;

/**
* @param businessLogic 接收 trace 增强后的 callback:onComplete / onError 会触发 finishRun
Expand All @@ -70,6 +74,7 @@ public void run(String question,

String traceId = IdUtil.getSnowflakeNextIdStr();
long startMillis = System.currentTimeMillis();
AtomicBoolean finished = new AtomicBoolean(false);
traceRecordService.startRun(RagTraceRunDO.builder()
.traceId(traceId)
.traceName(TRACE_NAME)
Expand All @@ -94,16 +99,27 @@ protected void onFirstContent() {

@Override
protected void onFinish(boolean success, Throwable error) {
finishRun(traceId, success, error, startMillis);
finishRunOnce(finished, traceId, taskId,
success ? STATUS_SUCCESS : STATUS_ERROR, error, startMillis);
}
};

RagTraceContext.setTraceId(traceId);
RagTraceContext.setTaskId(taskId);
try {
taskManager.bindCancellationObserver(taskId,
() -> finishRunOnce(finished, traceId, taskId, STATUS_CANCELLED, null, startMillis));
traceAwareCallback.onTraceStarted(traceId);
log.info("SSE stream started: traceId={}, taskId={}",
StreamTaskManager.safeCorrelationId(traceId), StreamTaskManager.safeCorrelationId(taskId));
if (taskManager.isCancelled(taskId)) {
return;
}
businessLogic.accept(traceAwareCallback);
} catch (Throwable ex) {
log.warn("执行流式对话失败(同步阶段),会话ID:{},任务ID:{}", conversationId, taskId, ex);
log.warn("SSE stream failed synchronously: traceId={}, taskId={}, errorType={}",
StreamTaskManager.safeCorrelationId(traceId), StreamTaskManager.safeCorrelationId(taskId),
ex.getClass().getSimpleName());
// 走 traceAwareCallback.onError 以复用其内部 CAS,避免与 pipeline 内已触发的终态重复收尾
try {
traceAwareCallback.onError(ex);
Expand Down Expand Up @@ -137,21 +153,34 @@ private void recordUserTtft(String traceId, Date runStartTime, long startMillis)
.build());
traceRecordService.finishNode(traceId, nodeId, STATUS_SUCCESS, null, new Date(now), durationMs);
} catch (Exception e) {
log.warn("写入 user-first-packet 节点失败,traceId:{}", traceId, e);
log.warn("Failed to record SSE first packet: traceId={}, errorType={}",
StreamTaskManager.safeCorrelationId(traceId), e.getClass().getSimpleName());
}
}

private void finishRun(String traceId, boolean success, Throwable error, long startMillis) {
private void finishRunOnce(AtomicBoolean finished,
String traceId,
String taskId,
String status,
Throwable error,
long startMillis) {
if (!finished.compareAndSet(false, true)) {
return;
}
try {
traceRecordService.finishRun(
traceId,
success ? STATUS_SUCCESS : STATUS_ERROR,
success ? null : truncateError(error),
status,
STATUS_ERROR.equals(status) ? truncateError(error) : null,
new Date(),
System.currentTimeMillis() - startMillis
);
log.info("SSE stream finished: traceId={}, taskId={}, status={}",
StreamTaskManager.safeCorrelationId(traceId),
StreamTaskManager.safeCorrelationId(taskId), status);
} catch (Exception e) {
log.warn("finishRun 失败,traceId:{}", traceId, e);
log.warn("Failed to finish SSE trace: traceId={}, errorType={}",
StreamTaskManager.safeCorrelationId(traceId), e.getClass().getSimpleName());
}
}

Expand All @@ -162,7 +191,8 @@ private void runWithoutTrace(String conversationId,
try {
businessLogic.accept(callback);
} catch (Throwable ex) {
log.warn("执行流式对话失败,会话ID:{},任务ID:{}", conversationId, taskId, ex);
log.warn("SSE stream failed without trace: taskId={}, errorType={}",
StreamTaskManager.safeCorrelationId(taskId), ex.getClass().getSimpleName());
callback.onError(ex);
}
}
Expand Down
Loading
Loading