94 lines
1.9 KiB
Go
94 lines
1.9 KiB
Go
package api
|
|
|
|
import (
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
|
|
"github.com/cloudwego/eino/adk"
|
|
"github.com/cloudwego/eino/schema"
|
|
"github.com/gin-gonic/gin"
|
|
)
|
|
|
|
type ChatHandler struct {
|
|
supervisor adk.Agent
|
|
}
|
|
|
|
func NewChatHandler(supervisor adk.Agent) *ChatHandler {
|
|
return &ChatHandler{supervisor: supervisor}
|
|
}
|
|
|
|
func (h *ChatHandler) Chat(c *gin.Context) {
|
|
var req ChatRequest
|
|
if err := c.ShouldBindJSON(&req); err != nil {
|
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
|
return
|
|
}
|
|
|
|
c.Header("Content-Type", "text/event-stream")
|
|
c.Header("Cache-Control", "no-cache")
|
|
c.Header("Connection", "keep-alive")
|
|
c.Header("X-Accel-Buffering", "no")
|
|
|
|
iter := h.supervisor.Run(c.Request.Context(), &adk.AgentInput{
|
|
Messages: []adk.Message{
|
|
schema.UserMessage(req.Message),
|
|
},
|
|
EnableStreaming: true,
|
|
})
|
|
|
|
c.Stream(func(w io.Writer) bool {
|
|
event, ok := iter.Next()
|
|
if !ok {
|
|
return false
|
|
}
|
|
|
|
if event.Err != nil {
|
|
writeSSE(w, "error", map[string]string{"error": event.Err.Error()})
|
|
return false
|
|
}
|
|
|
|
if event.Output == nil || event.Output.MessageOutput == nil {
|
|
return true
|
|
}
|
|
|
|
mv := event.Output.MessageOutput
|
|
|
|
if mv.IsStreaming && mv.MessageStream != nil {
|
|
stream := mv.MessageStream
|
|
for {
|
|
msg, err := stream.Recv()
|
|
if err == io.EOF {
|
|
break
|
|
}
|
|
if err != nil {
|
|
break
|
|
}
|
|
if msg.Content != "" {
|
|
writeSSE(w, "message", map[string]string{
|
|
"agent": event.AgentName,
|
|
"content": msg.Content,
|
|
"role": string(mv.Role),
|
|
})
|
|
}
|
|
}
|
|
} else if mv.Message != nil {
|
|
if mv.Message.Content != "" {
|
|
writeSSE(w, "message", map[string]string{
|
|
"agent": event.AgentName,
|
|
"content": mv.Message.Content,
|
|
"role": string(mv.Role),
|
|
})
|
|
}
|
|
}
|
|
|
|
return true
|
|
})
|
|
}
|
|
|
|
func writeSSE(w io.Writer, event string, data any) {
|
|
b, _ := json.Marshal(data)
|
|
fmt.Fprintf(w, "event: %s\ndata: %s\n\n", event, string(b))
|
|
}
|