Refactor message handling to support individual message lifecycle events

- Replaced group-related event types with message-specific types, enhancing clarity in event management.
- Updated stream handling functions to utilize new message start and end events, improving the granularity of message tracking.
- Refactored context management to integrate a context-scoped ID generator for unique message identifiers.
- Removed deprecated group handling methods and streamlined message processing logic.
- Updated tests and documentation to reflect changes in message handling and ensure proper functionality.
This commit is contained in:
Max 2025-11-26 19:03:47 +08:00
parent ef49827faf
commit 17b71ee4cc
12 changed files with 89 additions and 485 deletions

View file

@ -1,7 +1,6 @@
package handlers
import (
"fmt"
"time"
jsoniter "github.com/json-iterator/go"
@ -35,8 +34,8 @@ func DefaultStreamHandler(ctx *context.Context) message.StreamFunc {
case message.ChunkStreamStart:
return state.handleStreamStart(data)
case message.ChunkGroupStart:
return state.handleGroupStart(data)
case message.ChunkMessageStart:
return state.handleMessageStart(data)
case message.ChunkText:
return state.handleText(data)
@ -53,8 +52,8 @@ func DefaultStreamHandler(ctx *context.Context) message.StreamFunc {
case message.ChunkError:
return state.handleError(data)
case message.ChunkGroupEnd:
return state.handleGroupEnd(data)
case message.ChunkMessageEnd:
return state.handleMessageEnd(data)
case message.ChunkStreamEnd:
return state.handleStreamEnd(data)
@ -92,32 +91,32 @@ func (s *streamState) handleStreamStart(data []byte) int {
return 0
}
// handleGroupStart handles group start event
func (s *streamState) handleGroupStart(data []byte) int {
// Parse group start data first to get the group ID
// handleMessageStart handles message start event
func (s *streamState) handleMessageStart(data []byte) int {
// Parse message start data first to get the message ID
var startData message.EventMessageStartData
if err := jsoniter.Unmarshal(data, &startData); err != nil {
log.Error("Failed to unmarshal group start data: %v", err)
log.Error("Failed to unmarshal message start data: %v", err)
return 0
}
// Use the message ID from the start data, or generate one if not provided
groupID := startData.MessageID
if groupID == "" {
groupID = generateMessageID()
startData.MessageID = groupID
messageID := startData.MessageID
if messageID == "" {
messageID = s.ctx.IDGenerator.GenerateMessageID()
startData.MessageID = messageID
}
// Initialize group state with the correct group ID
// Initialize message state with the correct message ID
s.inGroup = true
s.currentGroupID = groupID
s.currentGroupID = messageID
s.buffer = []byte{}
s.chunkCount = 0
s.messageSeq = 0 // Reset message sequence for each group
s.messageSeq = 0 // Reset message sequence for each message
s.groupStartTime = time.Now()
// Send group_start event
msg := output.NewEventMessage(message.EventGroupStart, "Group started", startData)
// Send message_start event
msg := output.NewEventMessage(message.EventMessageStart, "Message started", startData)
s.ctx.Send(msg)
return 0 // Continue
@ -138,11 +137,11 @@ func (s *streamState) handleText(data []byte) int {
s.messageSeq++
// Send delta message
// - ChunkID: Sequential chunk ID (C1, C2, C3...) for this fragment
// - ChunkID: Unique chunk ID (C1, C2, C3...) for this fragment
// - MessageID: Same for all chunks of this logical message (frontend merges by message_id)
msg := &message.Message{
ChunkID: s.generateSequentialID(), // Sequential ID for this chunk
MessageID: s.currentGroupID, // Message ID for merging (all chunks share this)
ChunkID: s.ctx.IDGenerator.GenerateChunkID(), // Unique chunk ID
MessageID: s.currentGroupID, // Message ID for merging (all chunks share this)
Type: message.TypeText,
Delta: true,
Props: map[string]interface{}{
@ -173,11 +172,11 @@ func (s *streamState) handleThinking(data []byte) int {
s.messageSeq++
// Send delta message
// - ChunkID: Sequential chunk ID (C1, C2, C3...) for this fragment
// - ChunkID: Unique chunk ID (C1, C2, C3...) for this fragment
// - MessageID: Same for all chunks of this logical message (frontend merges by message_id)
msg := &message.Message{
ChunkID: s.generateSequentialID(), // Sequential ID for this chunk
MessageID: s.currentGroupID, // Message ID for merging (all chunks share this)
ChunkID: s.ctx.IDGenerator.GenerateChunkID(), // Unique chunk ID
MessageID: s.currentGroupID, // Message ID for merging (all chunks share this)
Type: message.TypeThinking,
Delta: true,
Props: map[string]interface{}{
@ -197,7 +196,8 @@ func (s *streamState) handleToolCall(data []byte) int {
// Tool calls are usually complete JSON objects
// Parse and send as tool_call message
msg := &message.Message{
MessageID: generateMessageID(),
ChunkID: s.ctx.IDGenerator.GenerateChunkID(),
MessageID: s.ctx.IDGenerator.GenerateMessageID(), // Tool call is a new message
Type: message.TypeToolCall,
Delta: true,
Props: map[string]interface{}{
@ -226,8 +226,8 @@ func (s *streamState) handleError(data []byte) int {
return 1 // Stop streaming on error
}
// handleGroupEnd handles group end event
func (s *streamState) handleGroupEnd(data []byte) int {
// handleMessageEnd handles message end event
func (s *streamState) handleMessageEnd(data []byte) int {
if !s.inGroup {
return 0
}
@ -254,8 +254,8 @@ func (s *streamState) handleGroupEnd(data []byte) int {
},
}
// Send group_end event
msg := output.NewEventMessage(message.EventGroupEnd, "Group completed", endData)
// Send message_end event
msg := output.NewEventMessage(message.EventMessageEnd, "Message completed", endData)
s.ctx.Send(msg)
// Reset state
@ -286,17 +286,3 @@ func (s *streamState) handleStreamEnd(data []byte) int {
s.ctx.Flush()
return 0 // Continue (stream will end naturally)
}
// generateSequentialID generates a sequential message ID for better readability
func (s *streamState) generateSequentialID() string {
// Format: 1, 2, 3, etc.
// This makes it easier for developers to track message order in logs
return fmt.Sprintf("%d", s.messageSeq)
}
// generateMessageID generates a unique message ID
func generateMessageID() string {
// TODO: Implement proper ID generation
// For now, use a simple approach
return output.GenerateID()
}

View file

@ -20,11 +20,11 @@ const (
ChunkError StreamChunkType = "error" // Error chunk
ChunkUnknown StreamChunkType = "unknown" // Unknown/unrecognized chunk type
// Lifecycle event types - stream and group boundaries
ChunkStreamStart StreamChunkType = "stream_start" // Stream begins (entire request starts)
ChunkStreamEnd StreamChunkType = "stream_end" // Stream ends (entire request completes)
ChunkGroupStart StreamChunkType = "group_start" // Message group begins (text/tool_call/thinking group starts)
ChunkGroupEnd StreamChunkType = "group_end" // Message group ends (text/tool_call/thinking group completes)
// Lifecycle event types - stream and message boundaries
ChunkStreamStart StreamChunkType = "stream_start" // Stream begins (entire request starts)
ChunkStreamEnd StreamChunkType = "stream_end" // Stream ends (entire request completes)
ChunkMessageStart StreamChunkType = "message_start" // Message begins (text/tool_call/thinking message starts)
ChunkMessageEnd StreamChunkType = "message_end" // Message ends (text/tool_call/thinking message completes)
)
// Writer is an alias for http.ResponseWriter interface used by an agent to construct a response.

View file

@ -7,6 +7,7 @@ import (
"time"
"github.com/yaoapp/gou/plan"
"github.com/yaoapp/yao/agent/output/message"
"github.com/yaoapp/yao/openapi/oauth/types"
)
@ -26,9 +27,10 @@ func newTestContextWithInterrupt(chatID, assistantID string) *Context {
UserAgent: "TestAgent/1.0",
IP: "127.0.0.1",
},
Referer: RefererAPI,
Accept: AcceptWebCUI,
Route: "/test/route",
Referer: RefererAPI,
Accept: AcceptWebCUI,
Route: "/test/route",
IDGenerator: message.NewIDGenerator(), // Initialize context-scoped ID generator
Metadata: map[string]interface{}{
"test": "context_metadata",
},

View file

@ -1,11 +1,8 @@
package context
import (
"time"
"github.com/yaoapp/gou/runtime/v8/bridge"
"github.com/yaoapp/yao/agent/output"
"github.com/yaoapp/yao/agent/output/message"
traceJsapi "github.com/yaoapp/yao/trace/jsapi"
"rogchap.com/v8go"
)
@ -50,9 +47,6 @@ func (ctx *Context) NewObject(v8ctx *v8go.Context) (*v8go.Value, error) {
// Set methods
jsObject.Set("Trace", ctx.traceMethod(v8ctx.Isolate()))
jsObject.Set("Send", ctx.sendMethod(v8ctx.Isolate()))
jsObject.Set("SendGroup", ctx.sendGroupMethod(v8ctx.Isolate()))
jsObject.Set("SendGroupStart", ctx.sendGroupStartMethod(v8ctx.Isolate()))
jsObject.Set("SendGroupEnd", ctx.sendGroupEndMethod(v8ctx.Isolate()))
// Create instance
instance, err := jsObject.NewInstance(v8ctx)
@ -211,188 +205,3 @@ func (ctx *Context) sendMethod(iso *v8go.Isolate) *v8go.FunctionTemplate {
// sendGroupMethod implements ctx.SendGroup(group)
// Usage: ctx.SendGroup({ id: "group1", messages: [...] })
// Automatically generates IDs, sends group_start/group_end events, and flushes output
func (ctx *Context) sendGroupMethod(iso *v8go.Isolate) *v8go.FunctionTemplate {
return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value {
v8ctx := info.Context()
args := info.Args()
if len(args) < 1 {
return bridge.JsException(v8ctx, "SendGroup requires a group argument")
}
// Parse group argument
group, err := parseGroup(v8ctx, args[0])
if err != nil {
return bridge.JsException(v8ctx, "invalid group: "+err.Error())
}
// Generate block ID if not provided
if group.ID == "" {
if ctx.IDGenerator != nil {
group.ID = ctx.IDGenerator.GenerateBlockID()
} else {
group.ID = output.GenerateID()
}
}
// Send group_start event
startTime := time.Now()
startEvent := output.NewEventMessage(
message.EventGroupStart,
"Group started",
message.EventMessageStartData{
MessageID: group.ID,
Type: "mixed", // Mixed types in group
Timestamp: startTime.UnixMilli(),
},
)
if err := ctx.Send(startEvent); err != nil {
return bridge.JsException(v8ctx, "Failed to send group_start event: "+err.Error())
}
if err := ctx.Flush(); err != nil {
return bridge.JsException(v8ctx, "Flush failed after group_start: "+err.Error())
}
// Generate MessageIDs for messages and set BlockID
for _, msg := range group.Messages {
if msg.MessageID == "" {
if ctx.IDGenerator != nil {
msg.MessageID = ctx.IDGenerator.GenerateMessageID()
} else {
msg.MessageID = output.GenerateID()
}
}
if msg.BlockID == "" {
msg.BlockID = group.ID
}
}
// Call ctx.SendGroup
if err := ctx.SendGroup(group); err != nil {
return bridge.JsException(v8ctx, "SendGroup failed: "+err.Error())
}
if err := ctx.Flush(); err != nil {
return bridge.JsException(v8ctx, "Flush failed after SendGroup: "+err.Error())
}
// Send group_end event
endEvent := output.NewEventMessage(
message.EventGroupEnd,
"Group completed",
message.EventMessageEndData{
MessageID: group.ID,
Type: "mixed",
Timestamp: time.Now().UnixMilli(),
DurationMs: time.Since(startTime).Milliseconds(),
ChunkCount: len(group.Messages),
Status: "completed",
},
)
if err := ctx.Send(endEvent); err != nil {
return bridge.JsException(v8ctx, "Failed to send group_end event: "+err.Error())
}
if err := ctx.Flush(); err != nil {
return bridge.JsException(v8ctx, "Flush failed after group_end: "+err.Error())
}
return v8go.Undefined(iso)
})
}
// sendGroupStartMethod implements ctx.SendGroupStart(type?, id?)
// Usage: const groupId = ctx.SendGroupStart() // type="mixed", auto-generate ID
// Usage: const groupId = ctx.SendGroupStart("text") // type="text", auto-generate ID
// Usage: const groupId = ctx.SendGroupStart("text", "my-group-id") // type="text", use provided ID
// Returns the group ID (generated or provided)
func (ctx *Context) sendGroupStartMethod(iso *v8go.Isolate) *v8go.FunctionTemplate {
return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value {
v8ctx := info.Context()
args := info.Args()
// Get type (default: "mixed")
groupType := "mixed"
if len(args) > 0 && args[0].IsString() {
groupType = args[0].String()
}
// Get or generate block ID
var groupID string
if len(args) > 1 && args[1].IsString() {
groupID = args[1].String()
} else {
if ctx.IDGenerator != nil {
groupID = ctx.IDGenerator.GenerateBlockID()
} else {
groupID = output.GenerateID()
}
}
// Send group_start event
startEvent := output.NewEventMessage(
message.EventGroupStart,
"Group started",
message.EventMessageStartData{
MessageID: groupID,
Type: groupType,
Timestamp: time.Now().UnixMilli(),
},
)
if err := ctx.Send(startEvent); err != nil {
return bridge.JsException(v8ctx, "Failed to send group_start event: "+err.Error())
}
if err := ctx.Flush(); err != nil {
return bridge.JsException(v8ctx, "Flush failed after group_start: "+err.Error())
}
// Return the group ID
groupIDVal, err := v8go.NewValue(iso, groupID)
if err != nil {
return bridge.JsException(v8ctx, "Failed to create return value: "+err.Error())
}
return groupIDVal
})
}
// sendGroupEndMethod implements ctx.SendGroupEnd(id, chunkCount?)
// Usage: ctx.SendGroupEnd(groupId)
// Usage: ctx.SendGroupEnd(groupId, 10) // With chunk count
func (ctx *Context) sendGroupEndMethod(iso *v8go.Isolate) *v8go.FunctionTemplate {
return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value {
v8ctx := info.Context()
args := info.Args()
// Group ID is required
if len(args) < 1 || !args[0].IsString() {
return bridge.JsException(v8ctx, "SendGroupEnd requires a group ID (string) as first argument")
}
groupID := args[0].String()
// Optional chunk count
chunkCount := 0
if len(args) > 1 && args[1].IsNumber() {
chunkCount = int(args[1].Integer())
}
// Send group_end event
endEvent := output.NewEventMessage(
message.EventGroupEnd,
"Group completed",
message.EventMessageEndData{
MessageID: groupID,
Type: "mixed",
Timestamp: time.Now().UnixMilli(),
DurationMs: 0, // Duration not tracked at this level
ChunkCount: chunkCount,
Status: "completed",
},
)
if err := ctx.Send(endEvent); err != nil {
return bridge.JsException(v8ctx, "Failed to send group_end event: "+err.Error())
}
if err := ctx.Flush(); err != nil {
return bridge.JsException(v8ctx, "Flush failed after group_end: "+err.Error())
}
return v8go.Undefined(iso)
})
}

View file

@ -8,6 +8,7 @@ import (
"github.com/stretchr/testify/assert"
v8 "github.com/yaoapp/gou/runtime/v8"
"github.com/yaoapp/yao/agent/output/message"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/test"
)
@ -52,6 +53,7 @@ func TestJsValueSend(t *testing.T) {
Accept: "standard",
Locale: "en",
Writer: newMockResponseWriter(),
IDGenerator: message.NewIDGenerator(),
}
// Test sending string shorthand
@ -108,106 +110,6 @@ func TestJsValueSend(t *testing.T) {
}
// TestJsValueSendGroup test the SendGroup method on Context
func TestJsValueSendGroup(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
cxt := &Context{
ChatID: "test-chat-id",
AssistantID: "test-assistant-id",
Context: context.Background(),
Accept: "standard",
Locale: "en",
Writer: newMockResponseWriter(),
}
res, err := v8.Call(v8.CallOptions{}, `
function test(ctx) {
try {
// Send message group
ctx.SendGroup({
id: "group_123",
messages: [
{
type: "text",
props: { content: "First message" }
},
{
type: "text",
props: { content: "Second message" }
},
{
type: "loading",
props: { message: "Processing..." }
}
],
metadata: {
timestamp: Date.now(),
sequence: 1
}
});
return { success: true };
} catch (error) {
return { success: false, error: error.message };
}
}`, cxt)
if err != nil {
t.Fatalf("Call failed: %v", err)
}
result, ok := res.(map[string]interface{})
if !ok {
t.Fatalf("Expected map result, got %T", res)
}
assert.Equal(t, true, result["success"], "SendGroup should succeed")
}
// TestJsValueSendGroupStartEnd test the SendGroupStart and SendGroupEnd methods
func TestJsValueSendGroupStartEnd(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
cxt := &Context{
ChatID: "test-chat-id",
AssistantID: "test-assistant-id",
Context: context.Background(),
Accept: "standard",
Locale: "en",
Writer: newMockResponseWriter(),
}
res, err := v8.Call(v8.CallOptions{}, `
function test(ctx) {
try {
// Start a group with auto-generated ID
const groupId = ctx.SendGroupStart("text");
// Send messages in the group
ctx.Send({ type: "text", props: { content: "Message 1" }, group_id: groupId });
ctx.Send({ type: "text", props: { content: "Message 2" }, group_id: groupId });
// End the group
ctx.SendGroupEnd(groupId, 2);
return { success: true, groupId: groupId };
} catch (error) {
return { success: false, error: error.message };
}
}`, cxt)
if err != nil {
t.Fatalf("Call failed: %v", err)
}
result, ok := res.(map[string]interface{})
if !ok {
t.Fatalf("Expected map result, got %T", res)
}
assert.Equal(t, true, result["success"], "SendGroupStart/End should succeed")
assert.NotEmpty(t, result["groupId"], "Should return group ID")
}
// TestJsValueSendDeltaUpdates test delta updates in Send
func TestJsValueSendDeltaUpdates(t *testing.T) {
@ -221,6 +123,7 @@ func TestJsValueSendDeltaUpdates(t *testing.T) {
Accept: "standard",
Locale: "en",
Writer: newMockResponseWriter(),
IDGenerator: message.NewIDGenerator(),
}
res, err := v8.Call(v8.CallOptions{}, `
@ -275,6 +178,7 @@ func TestJsValueSendMultipleTypes(t *testing.T) {
Accept: "standard",
Locale: "en",
Writer: newMockResponseWriter(),
IDGenerator: message.NewIDGenerator(),
}
res, err := v8.Call(v8.CallOptions{}, `
@ -357,6 +261,7 @@ func TestJsValueSendErrorHandling(t *testing.T) {
Accept: "standard",
Locale: "en",
Writer: newMockResponseWriter(),
IDGenerator: message.NewIDGenerator(),
}
// Test invalid argument - no arguments
@ -382,62 +287,6 @@ func TestJsValueSendErrorHandling(t *testing.T) {
}
// TestJsValueSendGroupErrorHandling test error handling in SendGroup
func TestJsValueSendGroupErrorHandling(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
cxt := &Context{
ChatID: "test-chat-id",
AssistantID: "test-assistant-id",
Context: context.Background(),
Accept: "standard",
Locale: "en",
Writer: newMockResponseWriter(),
}
// Test invalid argument - no arguments
res, err := v8.Call(v8.CallOptions{}, `
function test(ctx) {
try {
ctx.SendGroup();
return { success: true };
} catch (error) {
return { success: false, error: error.message };
}
}`, cxt)
if err != nil {
t.Fatalf("Call failed: %v", err)
}
result, ok := res.(map[string]interface{})
if !ok {
t.Fatalf("Expected map result, got %T", res)
}
assert.Equal(t, false, result["success"], "SendGroup without arguments should fail")
assert.Contains(t, result["error"], "SendGroup requires a group argument", "Error should mention missing group")
// Test invalid group - missing messages
res, err = v8.Call(v8.CallOptions{}, `
function test(ctx) {
try {
ctx.SendGroup({ id: "grp_1" });
return { success: true };
} catch (error) {
return { success: false, error: error.message };
}
}`, cxt)
if err != nil {
t.Fatalf("Call failed: %v", err)
}
result, ok = res.(map[string]interface{})
if !ok {
t.Fatalf("Expected map result, got %T", res)
}
assert.Equal(t, false, result["success"], "SendGroup without messages should fail")
}
// TestJsValueSendWithCUIAccept test Send with CUI accept types
func TestJsValueSendWithCUIAccept(t *testing.T) {
@ -483,67 +332,6 @@ func TestJsValueSendWithCUIAccept(t *testing.T) {
}
// TestJsValueSendGroupWithMetadata test SendGroup with various metadata
func TestJsValueSendGroupWithMetadata(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
cxt := &Context{
ChatID: "test-chat-id",
AssistantID: "test-assistant-id",
Context: context.Background(),
Accept: "standard",
Locale: "en",
Writer: newMockResponseWriter(),
}
res, err := v8.Call(v8.CallOptions{}, `
function test(ctx) {
try {
ctx.SendGroup({
id: "group_with_metadata",
messages: [
{
type: "text",
props: { content: "Message 1" },
metadata: {
timestamp: Date.now(),
sequence: 1,
trace_id: "trace_abc"
}
},
{
type: "text",
props: { content: "Message 2" },
metadata: {
timestamp: Date.now(),
sequence: 2,
trace_id: "trace_abc"
}
}
],
metadata: {
timestamp: Date.now(),
sequence: 1,
trace_id: "trace_abc"
}
});
return { success: true };
} catch (error) {
return { success: false, error: error.message };
}
}`, cxt)
if err != nil {
t.Fatalf("Call failed: %v", err)
}
result, ok := res.(map[string]interface{})
if !ok {
t.Fatalf("Expected map result, got %T", res)
}
assert.Equal(t, true, result["success"], "SendGroup with metadata should succeed")
}
// TestJsValueSendChainedCalls test chained Send calls
func TestJsValueSendChainedCalls(t *testing.T) {
@ -557,6 +345,7 @@ func TestJsValueSendChainedCalls(t *testing.T) {
Accept: "standard",
Locale: "en",
Writer: newMockResponseWriter(),
IDGenerator: message.NewIDGenerator(),
}
res, err := v8.Call(v8.CallOptions{}, `

View file

@ -9,6 +9,7 @@ import (
"github.com/stretchr/testify/assert"
v8 "github.com/yaoapp/gou/runtime/v8"
"github.com/yaoapp/gou/runtime/v8/bridge"
"github.com/yaoapp/yao/agent/output/message"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/openapi/oauth/types"
"github.com/yaoapp/yao/test"
@ -25,6 +26,7 @@ func TestJsValue(t *testing.T) {
ChatID: "ChatID-123456",
AssistantID: "AssistantID-1234",
Sid: "Sid-1234",
IDGenerator: message.NewIDGenerator(),
}
v8.RegisterFunction("testContextJsvalue", testContextJsvalueEmbed)
@ -93,6 +95,7 @@ func TestJsValueConcurrent(t *testing.T) {
ChatID: chatID,
AssistantID: assistantID,
Sid: sid,
IDGenerator: message.NewIDGenerator(),
}
res, err := v8.Call(v8.CallOptions{}, `
@ -151,6 +154,7 @@ func TestJsValueRegistrationAndCleanup(t *testing.T) {
ChatID: fmt.Sprintf("ChatID-%d", i),
AssistantID: fmt.Sprintf("AssistantID-%d", i),
Sid: fmt.Sprintf("Sid-%d", i),
IDGenerator: message.NewIDGenerator(),
}
_, err := v8.Call(v8.CallOptions{}, `
@ -441,7 +445,8 @@ func TestJsValueTrace(t *testing.T) {
Stack: &Stack{
TraceID: "test-trace-id",
},
Context: context.Background(),
Context: context.Background(),
IDGenerator: message.NewIDGenerator(),
}
res, err := v8.Call(v8.CallOptions{}, `

View file

@ -11,6 +11,7 @@ import (
"github.com/gin-gonic/gin"
"github.com/yaoapp/gou/plan"
"github.com/yaoapp/gou/store"
"github.com/yaoapp/yao/agent/output/message"
"github.com/yaoapp/yao/openapi/oauth/authorized"
)
@ -63,9 +64,10 @@ func GetCompletionRequest(c *gin.Context, cache store.Store) (*CompletionRequest
UserAgent: userAgent,
IP: clientIP,
},
Route: GetRoute(c, completionReq),
Metadata: GetMetadata(c, completionReq),
Skip: GetSkip(c, completionReq),
Route: GetRoute(c, completionReq),
Metadata: GetMetadata(c, completionReq),
Skip: GetSkip(c, completionReq),
IDGenerator: message.NewIDGenerator(), // Initialize context-scoped ID generator
}
// Initialize interrupt controller

View file

@ -4,6 +4,7 @@ import (
"testing"
"time"
"github.com/yaoapp/yao/agent/output/message"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/test"
)
@ -196,7 +197,9 @@ func TestEnterStack_RootCreation(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
ctx := &Context{}
ctx := &Context{
IDGenerator: message.NewIDGenerator(),
}
stack, traceID, done := EnterStack(ctx, "test-assistant", RefererAPI)
defer done()
@ -239,7 +242,9 @@ func TestEnterStack_ChildCreation(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
ctx := &Context{}
ctx := &Context{
IDGenerator: message.NewIDGenerator(),
}
// Create parent
parentStack, parentTraceID, parentDone := EnterStack(ctx, "parent-assistant", RefererAPI)
@ -282,7 +287,9 @@ func TestEnterStack_DoneCallback(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
ctx := &Context{}
ctx := &Context{
IDGenerator: message.NewIDGenerator(),
}
// Create parent
parentStack, _, parentDone := EnterStack(ctx, "parent-assistant", RefererAPI)
@ -321,7 +328,9 @@ func TestContextGetAllStacks(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
ctx := &Context{}
ctx := &Context{
IDGenerator: message.NewIDGenerator(),
}
// Create multiple stacks
_, _, done1 := EnterStack(ctx, "assistant1", RefererAPI)
@ -345,7 +354,9 @@ func TestContextGetStackByID(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
ctx := &Context{}
ctx := &Context{
IDGenerator: message.NewIDGenerator(),
}
stack, _, done := EnterStack(ctx, "test-assistant", RefererAPI)
defer done()
@ -372,7 +383,9 @@ func TestContextGetStacksByTraceID(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
ctx := &Context{}
ctx := &Context{
IDGenerator: message.NewIDGenerator(),
}
// Create parent and child (same trace ID)
_, traceID, done1 := EnterStack(ctx, "parent-assistant", RefererAPI)
@ -400,7 +413,9 @@ func TestContextGetRootStack(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
ctx := &Context{}
ctx := &Context{
IDGenerator: message.NewIDGenerator(),
}
// Create parent
parentStack, _, done1 := EnterStack(ctx, "parent-assistant", RefererAPI)

View file

@ -80,7 +80,7 @@ func TestDeepSeekR1StreamBasic(t *testing.T) {
}
// Track group_end events to verify type field
if chunkType == message.ChunkGroupEnd {
if chunkType == message.ChunkMessageEnd {
// Parse the group_end data to check the type field
var groupEndData struct {
GroupID string `json:"group_id"`

View file

@ -47,7 +47,7 @@ func (mt *messageTracker) startMessage(messageType message.StreamChunkType, hand
Timestamp: mt.startTime,
}
if startJSON, err := jsoniter.Marshal(startData); err == nil {
handler(message.ChunkGroupStart, startJSON)
handler(message.ChunkMessageStart, startJSON)
}
}
}
@ -79,7 +79,7 @@ func (mt *messageTracker) startToolCallMessage(toolCallInfo *message.EventToolCa
ToolCall: toolCallInfo,
}
if startJSON, err := jsoniter.Marshal(startData); err == nil {
handler(message.ChunkGroupStart, startJSON)
handler(message.ChunkMessageStart, startJSON)
}
}
}
@ -112,7 +112,7 @@ func (mt *messageTracker) endMessage(handler message.StreamFunc) {
endData.ToolCall = mt.toolCallInfo
}
if endJSON, err := jsoniter.Marshal(endData); err == nil {
handler(message.ChunkGroupEnd, endJSON)
handler(message.ChunkMessageEnd, endJSON)
}
}

View file

@ -1299,7 +1299,7 @@ func TestOpenAIStreamLifecycleEvents(t *testing.T) {
case message.ChunkStreamEnd:
t.Error("❌ LLM layer should NOT send stream_end (now sent at Agent level)")
case message.ChunkGroupStart:
case message.ChunkMessageStart:
groupStartReceived = true
var startData message.EventMessageStartData
if err := json.Unmarshal(data, &startData); err == nil {
@ -1311,7 +1311,7 @@ func TestOpenAIStreamLifecycleEvents(t *testing.T) {
t.Errorf("Failed to parse group_start data: %v", err)
}
case message.ChunkGroupEnd:
case message.ChunkMessageEnd:
groupEndReceived = true
var endData message.EventMessageEndData
if err := json.Unmarshal(data, &endData); err == nil {

View file

@ -110,10 +110,6 @@ const (
// Message level events (LLM layer - individual logical messages)
EventMessageStart = "message_start" // Message started event
EventMessageEnd = "message_end" // Message ended event
// Backward compatibility aliases (kept for transition period)
EventGroupStart = "group_start" // Alias for EventMessageStart
EventGroupEnd = "group_end" // Alias for EventMessageEnd
)
// Standard Props structures for built-in types
@ -243,11 +239,11 @@ const (
ChunkError StreamChunkType = "error" // Error chunk
ChunkUnknown StreamChunkType = "unknown" // Unknown/unrecognized chunk type
// Lifecycle event types - stream and group boundaries
ChunkStreamStart StreamChunkType = "stream_start" // Stream begins (entire request starts)
ChunkStreamEnd StreamChunkType = "stream_end" // Stream ends (entire request completes)
ChunkGroupStart StreamChunkType = "group_start" // Message group begins (text/tool_call/thinking group starts)
ChunkGroupEnd StreamChunkType = "group_end" // Message group ends (text/tool_call/thinking group completes)
// Lifecycle event types - stream and message boundaries
ChunkStreamStart StreamChunkType = "stream_start" // Stream begins (entire request starts)
ChunkStreamEnd StreamChunkType = "stream_end" // Stream ends (entire request completes)
ChunkMessageStart StreamChunkType = "message_start" // Message begins (text/tool_call/thinking message starts)
ChunkMessageEnd StreamChunkType = "message_end" // Message ends (text/tool_call/thinking message completes)
)
// StreamFunc the streaming function callback