164 lines
3.9 KiB
Go
164 lines
3.9 KiB
Go
package handler
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"io"
|
|
"net/http"
|
|
"resume-platform/internal/service/chat"
|
|
"resume-platform/pkg/logger"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/gorilla/websocket"
|
|
)
|
|
|
|
var upgrader = websocket.Upgrader{
|
|
ReadBufferSize: 4096,
|
|
WriteBufferSize: 4096,
|
|
CheckOrigin: func(r *http.Request) bool {
|
|
return true
|
|
},
|
|
}
|
|
|
|
type IMHandler struct {
|
|
imService *chat.IMService
|
|
}
|
|
|
|
func NewIMHandler(imService *chat.IMService) *IMHandler {
|
|
return &IMHandler{imService: imService}
|
|
}
|
|
|
|
func (h *IMHandler) WebSocket(c *gin.Context) {
|
|
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
|
if err != nil {
|
|
logger.Errorf("Failed to upgrade connection: %v", err)
|
|
return
|
|
}
|
|
defer conn.Close()
|
|
|
|
ctx := c.Request.Context()
|
|
logger.CtxInfof(ctx, "WebSocket connection opened")
|
|
|
|
for {
|
|
_, msg, err := conn.ReadMessage()
|
|
if err != nil {
|
|
if websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) {
|
|
logger.CtxInfof(ctx, "WebSocket connection closed normally")
|
|
} else {
|
|
logger.CtxErrorf(ctx, "WebSocket read error: %v", err)
|
|
}
|
|
break
|
|
}
|
|
|
|
var request struct {
|
|
Message string `json:"message"`
|
|
}
|
|
if err := json.Unmarshal(msg, &request); err != nil {
|
|
logger.CtxErrorf(ctx, "Failed to parse message: %v", err)
|
|
continue
|
|
}
|
|
|
|
if request.Message == "" {
|
|
continue
|
|
}
|
|
|
|
go h.handleMessage(ctx, conn, request.Message)
|
|
}
|
|
}
|
|
|
|
func (h *IMHandler) handleMessage(ctx context.Context, conn *websocket.Conn, message string) {
|
|
logger.CtxInfof(ctx, "Received message: %s", message)
|
|
|
|
agent := h.imService.GetResumeAgent()
|
|
var result string
|
|
var err error
|
|
if agent != nil {
|
|
logger.CtxInfof(ctx, "Using Eino Resume Agent for request")
|
|
result, err = h.imService.GenerateWithAgent(ctx, message)
|
|
} else {
|
|
logger.CtxInfof(ctx, "Using AI Provider for request")
|
|
result, err = h.imService.GetAIProvider().Generate(ctx, message)
|
|
}
|
|
|
|
if err != nil {
|
|
logger.CtxErrorf(ctx, "AI generate failed: %v", err)
|
|
response := map[string]interface{}{
|
|
"type": "error",
|
|
"message": err.Error(),
|
|
}
|
|
if err := conn.WriteJSON(response); err != nil {
|
|
logger.CtxErrorf(ctx, "Failed to send error response: %v", err)
|
|
}
|
|
return
|
|
}
|
|
|
|
response := map[string]interface{}{
|
|
"type": "message",
|
|
"content": result,
|
|
}
|
|
|
|
if err := conn.WriteJSON(response); err != nil {
|
|
logger.CtxErrorf(ctx, "Failed to send response: %v", err)
|
|
}
|
|
}
|
|
|
|
func (h *IMHandler) Chat(c *gin.Context) {
|
|
var request struct {
|
|
Message string `json:"message"`
|
|
}
|
|
if err := c.ShouldBindJSON(&request); err != nil {
|
|
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid request"})
|
|
return
|
|
}
|
|
|
|
if request.Message == "" {
|
|
c.JSON(http.StatusBadRequest, gin.H{"error": "message is required"})
|
|
return
|
|
}
|
|
|
|
ctx := c.Request.Context()
|
|
logger.CtxInfof(ctx, "Chat request: %s", request.Message)
|
|
|
|
result, err := h.imService.GetAIProvider().Generate(ctx, request.Message)
|
|
if err != nil {
|
|
logger.CtxErrorf(ctx, "AI generate failed: %v", err)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
|
return
|
|
}
|
|
|
|
c.JSON(http.StatusOK, gin.H{"response": result})
|
|
}
|
|
|
|
func (h *IMHandler) ChatStream(c *gin.Context) {
|
|
var request struct {
|
|
Message string `json:"message"`
|
|
}
|
|
if err := c.ShouldBindJSON(&request); err != nil {
|
|
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid request"})
|
|
return
|
|
}
|
|
|
|
if request.Message == "" {
|
|
c.JSON(http.StatusBadRequest, gin.H{"error": "message is required"})
|
|
return
|
|
}
|
|
|
|
ctx := c.Request.Context()
|
|
logger.CtxInfof(ctx, "Chat stream request: %s", request.Message)
|
|
|
|
c.Stream(func(w io.Writer) bool {
|
|
result, err := h.imService.GetAIProvider().Generate(ctx, request.Message)
|
|
if err != nil {
|
|
logger.CtxErrorf(ctx, "AI generate failed: %v", err)
|
|
c.SSEvent("error", gin.H{"message": err.Error()})
|
|
return false
|
|
}
|
|
|
|
for _, char := range result {
|
|
c.SSEvent("message", string(char))
|
|
}
|
|
c.SSEvent("done", gin.H{})
|
|
return false
|
|
})
|
|
}
|