diff --git a/internal/ai/runtime/graphs/handoff_graph.go b/internal/ai/runtime/graphs/handoff_graph.go index 8b6a21e..2b538ab 100644 --- a/internal/ai/runtime/graphs/handoff_graph.go +++ b/internal/ai/runtime/graphs/handoff_graph.go @@ -8,6 +8,7 @@ import ( "cs-agent/internal/ai/runtime/tooling" "cs-agent/internal/models" + "cs-agent/internal/pkg/tracex" "cs-agent/internal/services" componenttool "github.com/cloudwego/eino/components/tool" @@ -51,7 +52,8 @@ func (g *HandoffGraph) Run(ctx context.Context, argumentsInJSON string) (string, if err != nil { return "", err } - handled, err := services.ConversationService.TryOffHoursHandoffByAI(g.conversation.ID, g.aiAgent, reason) + requestID := tracex.RequestIDFromContext(ctx) + handled, err := services.ConversationService.TryOffHoursHandoffByAIWithRequestID(g.conversation.ID, g.aiAgent, reason, requestID) if err != nil || handled { if handled && err == nil { return tooling.MarshalToolResult(tooling.ToolResult{ @@ -91,7 +93,7 @@ func (g *HandoffGraph) Run(ctx context.Context, argumentsInJSON string) (string, } switch parseHandoffDecision(resumeText) { case ConfirmationDecisionConfirm: - if err := services.ConversationService.HandoffByAI(g.conversation.ID, g.aiAgent, state.Reason); err != nil { + if err := services.ConversationService.HandoffByAIWithRequestID(g.conversation.ID, g.aiAgent, state.Reason, tracex.RequestIDFromContext(ctx)); err != nil { return "", err } // ConversationService sends the customer-visible handoff notice according to the dispatch decision. diff --git a/internal/ai/runtime/reply_commit_service.go b/internal/ai/runtime/reply_commit_service.go index cef8cb0..26c78b3 100644 --- a/internal/ai/runtime/reply_commit_service.go +++ b/internal/ai/runtime/reply_commit_service.go @@ -36,7 +36,7 @@ func (s *replyCommitService) SendAIReply(input replyCommitInput) (*models.Messag return nil, nil } commitStartedAt := time.Now() - replyMessage, err := svc.MessageService.SendAIMessage( + replyMessage, err := svc.MessageService.SendAIMessageWithRequestID( input.Conversation.ID, input.AIAgent.ID, fmt.Sprintf("%s_%d", strings.TrimSpace(input.ClientPrefix), input.Message.ID), @@ -44,6 +44,7 @@ func (s *replyCommitService) SendAIReply(input replyCommitInput) (*models.Messag replyText, "", s.buildAIPrincipal(input.AIAgent), + input.Message.RequestID, ) if input.Trace != nil { input.Trace.CommitMs = time.Since(commitStartedAt).Milliseconds() diff --git a/internal/ai/runtime/reply_runlog_service.go b/internal/ai/runtime/reply_runlog_service.go index c5e2ebd..e99591d 100644 --- a/internal/ai/runtime/reply_runlog_service.go +++ b/internal/ai/runtime/reply_runlog_service.go @@ -41,6 +41,7 @@ func (s *replyRunLogService) Write(input replyRunLogInput) { logItem := &models.AgentRunLog{ ConversationID: input.Conversation.ID, MessageID: input.Message.ID, + RequestID: input.Message.RequestID, AIAgentID: input.AIAgent.ID, AIConfigID: input.AIAgent.AIConfigID, UserMessage: strings.TrimSpace(input.Question), @@ -66,6 +67,7 @@ func (s *replyRunLogService) Write(input replyRunLogInput) { } if err := svc.AgentRunLogService.Create(logItem); err != nil { slog.Warn("create agent run log failed", + "requestId", input.Message.RequestID, "message_id", input.Message.ID, "conversation_id", logItem.ConversationID, "ai_agent_id", input.AIAgent.ID, @@ -245,6 +247,9 @@ func extractGraphToolTrace(summary *applicationruntime.Summary) string { } func firstToolSearchTargetToolCode(summary *applicationruntime.Summary) string { + if summary == nil { + return "" + } trace := parseRuntimeTraceData(summary.TraceData) for _, item := range trace.ToolSearch.Items { toolCode := strings.TrimSpace(item.TargetToolCode) @@ -262,6 +267,9 @@ func firstToolSearchTargetToolCode(summary *applicationruntime.Summary) string { } func firstGraphToolCode(summary *applicationruntime.Summary) string { + if summary == nil { + return "" + } trace := parseRuntimeTraceData(summary.TraceData) for _, item := range trace.GraphTools.Items { toolCode := strings.TrimSpace(item.ToolCode) @@ -273,6 +281,9 @@ func firstGraphToolCode(summary *applicationruntime.Summary) string { } func extractHandoffReason(summary *applicationruntime.Summary) string { + if summary == nil { + return "" + } trace := parseRuntimeTraceData(summary.TraceData) for _, item := range trace.GraphTools.Items { if strings.TrimSpace(item.ToolCode) != toolx.GraphHandoffConversation.Code { @@ -291,6 +302,9 @@ func extractHandoffReason(summary *applicationruntime.Summary) string { } func graphPlanReason(summary *applicationruntime.Summary) string { + if summary == nil { + return "" + } trace := parseRuntimeTraceData(summary.TraceData) for _, item := range trace.GraphTools.Items { toolCode := strings.TrimSpace(item.ToolCode) diff --git a/internal/ai/runtime/reply_runlog_service_test.go b/internal/ai/runtime/reply_runlog_service_test.go new file mode 100644 index 0000000..3b8c064 --- /dev/null +++ b/internal/ai/runtime/reply_runlog_service_test.go @@ -0,0 +1,57 @@ +package runtime + +import ( + "strings" + "testing" + "time" + + "cs-agent/internal/models" + "cs-agent/internal/pkg/enums" + + "github.com/glebarez/sqlite" + "github.com/mlogclub/simple/sqls" + "gorm.io/gorm" + "gorm.io/gorm/schema" +) + +func TestReplyRunLogStoresRequestID(t *testing.T) { + dbName := "reply_runlog_trace_test_" + strings.NewReplacer("/", "_").Replace(t.Name()) + db, err := gorm.Open(sqlite.Open("file:"+dbName+"?mode=memory&cache=shared"), &gorm.Config{ + NamingStrategy: schema.NamingStrategy{ + TablePrefix: "t_", + SingularTable: true, + }, + }) + if err != nil { + t.Fatalf("open sqlite db: %v", err) + } + sqlDB, err := db.DB() + if err != nil { + t.Fatalf("get sqlite db: %v", err) + } + t.Cleanup(func() { + if err := sqlDB.Close(); err != nil { + t.Fatalf("close sqlite db: %v", err) + } + }) + if err := db.AutoMigrate(&models.AgentRunLog{}); err != nil { + t.Fatalf("auto migrate: %v", err) + } + sqls.SetDB(db) + + newReplyRunLogService().Write(replyRunLogInput{ + StartedAt: time.Now(), + Message: models.Message{ID: 22, RequestID: "trace-123", SenderType: enums.IMSenderTypeCustomer, Content: "hello"}, + Conversation: models.Conversation{ID: 11}, + AIAgent: models.AIAgent{ID: 33, AIConfigID: 44}, + Question: "hello", + }) + + var item models.AgentRunLog + if err := db.First(&item).Error; err != nil { + t.Fatalf("find run log: %v", err) + } + if item.RequestID != "trace-123" { + t.Fatalf("RequestID=%q want %q", item.RequestID, "trace-123") + } +} diff --git a/internal/ai/runtime/reply_trigger_service.go b/internal/ai/runtime/reply_trigger_service.go index cf57a21..a368cdb 100644 --- a/internal/ai/runtime/reply_trigger_service.go +++ b/internal/ai/runtime/reply_trigger_service.go @@ -9,6 +9,7 @@ import ( applicationruntime "cs-agent/internal/ai/application/runtime" "cs-agent/internal/models" "cs-agent/internal/pkg/enums" + "cs-agent/internal/pkg/tracex" svc "cs-agent/internal/services" ) @@ -30,10 +31,11 @@ func (s *aiReplyService) TriggerReplyAsync(conversation models.Conversation, mes } startedAt := time.Now() timeout := s.resolveReplyTimeout(*aiAgent) - ctx, cancel := context.WithTimeout(context.Background(), timeout) + ctx, cancel := context.WithTimeout(tracex.ContextWithRequestID(context.Background(), message.RequestID), timeout) defer cancel() if err := s.TriggerReply(ctx, conversation, message, *aiAgent); err != nil { slog.Error("failed to trigger ai reply", + "requestId", message.RequestID, "message_id", message.ID, "timeout_ms", timeout.Milliseconds(), "elapsed_ms", time.Since(startedAt).Milliseconds(), diff --git a/internal/bootstrap/routes.go b/internal/bootstrap/routes.go index 683c06c..162ff73 100644 --- a/internal/bootstrap/routes.go +++ b/internal/bootstrap/routes.go @@ -11,6 +11,7 @@ import ( func registerApiAuthRoutes(group *gin.RouterGroup) { group.POST("/login", api.Login) group.POST("/logout", api.Logout) + group.GET("/options", api.AuthOptions) group.GET("/profile", api.Profile) group.GET("/wxwork_callback", api.WxWorkCallback) group.POST("/wxwork_exchange", api.WxWorkExchange) diff --git a/internal/bootstrap/server.go b/internal/bootstrap/server.go index fe35979..81cfe55 100644 --- a/internal/bootstrap/server.go +++ b/internal/bootstrap/server.go @@ -13,6 +13,7 @@ import ( "cs-agent/internal/pkg/ginx" "cs-agent/internal/pkg/httpx" "cs-agent/internal/pkg/i18nx" + "cs-agent/internal/pkg/tracex" "cs-agent/internal/services" webspa "cs-agent/web" @@ -29,6 +30,7 @@ func NewServer() (*gin.Engine, error) { printBanner() app := gin.New() + app.Use(requestIDMiddleware()) app.Use(corsMiddleware()) app.Use(gin.Recovery()) app.Use(requestLogMiddleware()) @@ -103,14 +105,27 @@ func corsMiddleware() gin.HandlerFunc { } } +func requestIDMiddleware() gin.HandlerFunc { + return func(ctx *gin.Context) { + requestID := tracex.EnsureRequestID(ctx.GetHeader(tracex.RequestIDHeader)) + ctx.Set(tracex.GinRequestIDKey, requestID) + if requestID != "" { + ctx.Header(tracex.RequestIDHeader, requestID) + } + ctx.Next() + } +} + func requestLogMiddleware() gin.HandlerFunc { return func(ctx *gin.Context) { start := time.Now() path := ctx.Request.URL.Path method := ctx.Request.Method + requestID, _ := ctx.Get(tracex.GinRequestIDKey) ctx.Next() slog.Info("http request", + "requestId", requestID, "method", method, "path", path, "status", ctx.Writer.Status(), diff --git a/internal/bootstrap/server_route_test.go b/internal/bootstrap/server_route_test.go index 1eb2b42..a4910dd 100644 --- a/internal/bootstrap/server_route_test.go +++ b/internal/bootstrap/server_route_test.go @@ -1,6 +1,7 @@ package bootstrap import ( + "encoding/json" "net/http" "net/http/httptest" "strings" @@ -49,6 +50,59 @@ func TestNewServerRegistersGinRoutes(t *testing.T) { } } +func TestNewServerExposesPublicAuthOptions(t *testing.T) { + config.SetCurrent(&config.Config{ + Storage: config.StorageConfig{ + Local: config.LocalStorageConfig{ + Root: "storage", + BaseURL: "/storage", + }, + }, + WxWork: config.WxWorkConfig{ + Enabled: true, + }, + OIDC: config.OIDCConfig{ + Enabled: false, + ClientSecret: "must-not-leak", + }, + }) + + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + rec := httptest.NewRecorder() + app.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/auth/options", nil)) + + if rec.Code != http.StatusOK { + t.Fatalf("status=%d want %d", rec.Code, http.StatusOK) + } + + var body struct { + Success bool `json:"success"` + Data struct { + WxWorkEnabled bool `json:"wxworkEnabled"` + OIDCEnabled bool `json:"oidcEnabled"` + } `json:"data"` + } + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("unmarshal response: %v", err) + } + if !body.Success { + t.Fatalf("success=false, body=%s", rec.Body.String()) + } + if !body.Data.WxWorkEnabled { + t.Fatalf("wxworkEnabled=false want true") + } + if body.Data.OIDCEnabled { + t.Fatalf("oidcEnabled=true want false") + } + if strings.Contains(rec.Body.String(), "must-not-leak") { + t.Fatalf("response leaked sensitive OIDC config: %s", rec.Body.String()) + } +} + func TestNewServerSeparatesAPIStaticAndSPA(t *testing.T) { config.SetCurrent(&config.Config{ Storage: config.StorageConfig{ @@ -159,3 +213,51 @@ func TestNewServerRejectsUnconfiguredCORSOrigin(t *testing.T) { t.Fatalf("Access-Control-Allow-Origin=%q want empty", got) } } + +func TestNewServerEchoesRequestID(t *testing.T) { + config.SetCurrent(&config.Config{ + Storage: config.StorageConfig{ + Local: config.LocalStorageConfig{ + Root: "storage", + BaseURL: "/storage", + }, + }, + }) + + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/not-exists", nil) + req.Header.Set("X-Request-Id", "trace-123") + app.ServeHTTP(rec, req) + + if got := rec.Header().Get("X-Request-Id"); got != "trace-123" { + t.Fatalf("X-Request-Id=%q want %q", got, "trace-123") + } +} + +func TestNewServerGeneratesRequestID(t *testing.T) { + config.SetCurrent(&config.Config{ + Storage: config.StorageConfig{ + Local: config.LocalStorageConfig{ + Root: "storage", + BaseURL: "/storage", + }, + }, + }) + + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + rec := httptest.NewRecorder() + app.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/not-exists", nil)) + + if got := rec.Header().Get("X-Request-Id"); got == "" { + t.Fatalf("X-Request-Id should be generated") + } +} diff --git a/internal/builders/agent_run_log_builder.go b/internal/builders/agent_run_log_builder.go index f6b0720..cbaf3c4 100644 --- a/internal/builders/agent_run_log_builder.go +++ b/internal/builders/agent_run_log_builder.go @@ -17,6 +17,7 @@ func BuildAgentRunLog(item *models.AgentRunLog) response.AgentRunLogResponse { ID: item.ID, ConversationID: item.ConversationID, MessageID: item.MessageID, + RequestID: item.RequestID, AIAgentID: item.AIAgentID, AIConfigID: item.AIConfigID, UserMessage: item.UserMessage, diff --git a/internal/builders/conversation_builder.go b/internal/builders/conversation_builder.go index 0e54d3f..2f42255 100644 --- a/internal/builders/conversation_builder.go +++ b/internal/builders/conversation_builder.go @@ -149,6 +149,7 @@ func BuildMessageWithReadStatesAndLocale(item *models.Message, agentReadState, c ret := response.MessageResponse{ ID: item.ID, ConversationID: item.ConversationID, + RequestID: item.RequestID, ClientMsgID: item.ClientMsgID, SenderType: item.SenderType, SenderID: item.SenderID, diff --git a/internal/handlers/api/auth_handler.go b/internal/handlers/api/auth_handler.go index a552dc8..2f1d06d 100644 --- a/internal/handlers/api/auth_handler.go +++ b/internal/handlers/api/auth_handler.go @@ -3,6 +3,7 @@ package api import ( "cs-agent/internal/pkg/config" "cs-agent/internal/pkg/dto/request" + "cs-agent/internal/pkg/dto/response" "cs-agent/internal/pkg/httpx" "cs-agent/internal/pkg/httpx/params" "cs-agent/internal/services" @@ -29,6 +30,14 @@ func Login(ctx *gin.Context) { httpx.WriteJSON(ctx, ret) } +func AuthOptions(ctx *gin.Context) { + cfg := config.Current() + httpx.WriteJSON(ctx, &response.AuthOptionsResponse{ + WxWorkEnabled: cfg.WxWork.Enabled, + OIDCEnabled: cfg.OIDC.Enabled, + }) +} + func WxWorkLogin(ctx *gin.Context) { loginURL, err := services.WxWorkLoginService.BuildWxWorkLoginURL(ctx.Query("next")) if err != nil { diff --git a/internal/handlers/api/message_handler.go b/internal/handlers/api/message_handler.go index d3581b6..a57fc75 100644 --- a/internal/handlers/api/message_handler.go +++ b/internal/handlers/api/message_handler.go @@ -73,7 +73,7 @@ func MessagePostSend(ctx *gin.Context) { return } - item, err := services.MessageService.SendCustomerMessage(req.ConversationID, req.ClientMsgID, req.MessageType, req.Content, req.Payload, *external) + item, err := services.MessageService.SendCustomerMessageWithRequestID(req.ConversationID, req.ClientMsgID, req.MessageType, req.Content, req.Payload, *external, httpx.GetRequestID(ctx)) if err != nil { httpx.WriteJSON(ctx, err) return diff --git a/internal/handlers/dashboard/agent_run_log_handler.go b/internal/handlers/dashboard/agent_run_log_handler.go index 80c130d..55d1fbe 100644 --- a/internal/handlers/dashboard/agent_run_log_handler.go +++ b/internal/handlers/dashboard/agent_run_log_handler.go @@ -22,6 +22,7 @@ func AgentRunLogAnyList(ctx *gin.Context) { cnd := params.NewPagedSqlCnd(ctx, params.QueryFilter{ParamName: "conversationId"}, params.QueryFilter{ParamName: "messageId"}, + params.QueryFilter{ParamName: "requestId"}, params.QueryFilter{ParamName: "aiAgentId"}, params.QueryFilter{ParamName: "plannedAction"}, params.QueryFilter{ParamName: "plannedSkillCode", Op: params.Like}, diff --git a/internal/handlers/dashboard/conversation_handler.go b/internal/handlers/dashboard/conversation_handler.go index 40adae8..2b66955 100644 --- a/internal/handlers/dashboard/conversation_handler.go +++ b/internal/handlers/dashboard/conversation_handler.go @@ -255,7 +255,7 @@ func ConversationPostSend_message(ctx *gin.Context) { httpx.WriteJSON(ctx, err) return } - item, err := services.MessageService.SendAgentMessage(req.ConversationID, 0, req.ClientMsgID, req.MessageType, req.Content, req.Payload, operator) + item, err := services.MessageService.SendAgentMessageWithRequestID(req.ConversationID, 0, req.ClientMsgID, req.MessageType, req.Content, req.Payload, operator, httpx.GetRequestID(ctx)) if err != nil { httpx.WriteJSON(ctx, err) return diff --git a/internal/models/models.go b/internal/models/models.go index 019d07e..2b64d0c 100644 --- a/internal/models/models.go +++ b/internal/models/models.go @@ -385,6 +385,7 @@ type ConversationReadState struct { type Message struct { ID int64 `gorm:"primaryKey;autoIncrement"` ConversationID int64 `gorm:"type:bigint;not null;index;uniqueIndex:uk_conversation_seq;uniqueIndex:uk_conversation_client_msg"` + RequestID string `gorm:"type:varchar(128);not null;default:'';index"` ClientMsgID string `gorm:"type:varchar(128);not null;default:'';uniqueIndex:uk_conversation_client_msg"` SenderType enums.IMSenderType `gorm:"type:varchar(30);not null;default:'';index"` SenderID int64 `gorm:"type:bigint;not null;default:0;index"` @@ -552,6 +553,7 @@ type Channel struct { type ConversationEventLog struct { ID int64 `gorm:"primaryKey;autoIncrement"` ConversationID int64 `gorm:"type:bigint;not null;index"` + RequestID string `gorm:"type:varchar(128);not null;default:'';index"` EventType enums.IMEventType `gorm:"type:varchar(50);not null;default:'';index"` OperatorType enums.IMSenderType `gorm:"type:varchar(30);not null;default:'';index"` OperatorID int64 `gorm:"type:bigint;not null;default:0;index"` @@ -839,6 +841,7 @@ type AgentRunLog struct { ID int64 `gorm:"primaryKey;autoIncrement"` ConversationID int64 `gorm:"type:bigint;not null;default:0;index"` MessageID int64 `gorm:"type:bigint;not null;default:0;index"` + RequestID string `gorm:"type:varchar(128);not null;default:'';index"` AIAgentID int64 `gorm:"type:bigint;not null;default:0;index"` AIConfigID int64 `gorm:"type:bigint;not null;default:0;index"` UserMessage string `gorm:"type:longtext"` diff --git a/internal/pkg/dto/response/auth_response.go b/internal/pkg/dto/response/auth_response.go index e0c2a1d..9fe1261 100644 --- a/internal/pkg/dto/response/auth_response.go +++ b/internal/pkg/dto/response/auth_response.go @@ -18,3 +18,8 @@ type LoginResponse struct { Permissions []string `json:"permissions"` Roles []string `json:"roles"` } + +type AuthOptionsResponse struct { + WxWorkEnabled bool `json:"wxworkEnabled"` + OIDCEnabled bool `json:"oidcEnabled"` +} diff --git a/internal/pkg/dto/response/message_response.go b/internal/pkg/dto/response/message_response.go index e1b0c4c..11b3f18 100644 --- a/internal/pkg/dto/response/message_response.go +++ b/internal/pkg/dto/response/message_response.go @@ -5,6 +5,7 @@ import "cs-agent/internal/pkg/enums" type MessageResponse struct { ID int64 `json:"id"` ConversationID int64 `json:"conversationId"` + RequestID string `json:"requestId,omitempty"` ClientMsgID string `json:"clientMsgId,omitempty"` SenderType enums.IMSenderType `json:"senderType"` SenderID int64 `json:"senderId"` diff --git a/internal/pkg/dto/response/skill_response.go b/internal/pkg/dto/response/skill_response.go index 253d36f..6135439 100644 --- a/internal/pkg/dto/response/skill_response.go +++ b/internal/pkg/dto/response/skill_response.go @@ -44,6 +44,7 @@ type AgentRunLogResponse struct { ID int64 `json:"id"` ConversationID int64 `json:"conversationId"` MessageID int64 `json:"messageId"` + RequestID string `json:"requestId"` AIAgentID int64 `json:"aiAgentId"` AIConfigID int64 `json:"aiConfigId"` UserMessage string `json:"userMessage"` diff --git a/internal/pkg/httpx/context.go b/internal/pkg/httpx/context.go index 33278d8..b0bacd1 100644 --- a/internal/pkg/httpx/context.go +++ b/internal/pkg/httpx/context.go @@ -3,6 +3,7 @@ package httpx import ( "cs-agent/internal/pkg/httpx/params" "cs-agent/internal/pkg/openidentity" + "cs-agent/internal/pkg/tracex" "github.com/gin-gonic/gin" "github.com/mlogclub/simple/common/strs" @@ -31,3 +32,15 @@ func GetChannelID(ctx *gin.Context) string { } return "" } + +func GetRequestID(ctx *gin.Context) string { + if ctx == nil { + return "" + } + if value, ok := ctx.Get(tracex.GinRequestIDKey); ok { + if requestID, ok := value.(string); ok { + return tracex.NormalizeRequestID(requestID) + } + } + return tracex.NormalizeRequestID(ctx.GetHeader(tracex.RequestIDHeader)) +} diff --git a/internal/pkg/tracex/request_id.go b/internal/pkg/tracex/request_id.go new file mode 100644 index 0000000..30defc4 --- /dev/null +++ b/internal/pkg/tracex/request_id.go @@ -0,0 +1,57 @@ +package tracex + +import ( + "context" + "crypto/rand" + "encoding/hex" + "strings" +) + +const ( + RequestIDHeader = "X-Request-Id" + GinRequestIDKey = "requestId" +) + +type requestIDContextKey struct{} + +func NormalizeRequestID(value string) string { + value = strings.TrimSpace(value) + if value == "" || len(value) > 128 { + return "" + } + for _, r := range value { + if r < 33 || r > 126 { + return "" + } + } + return value +} + +func EnsureRequestID(value string) string { + if normalized := NormalizeRequestID(value); normalized != "" { + return normalized + } + var b [16]byte + if _, err := rand.Read(b[:]); err != nil { + return "" + } + return hex.EncodeToString(b[:]) +} + +func ContextWithRequestID(ctx context.Context, requestID string) context.Context { + requestID = NormalizeRequestID(requestID) + if requestID == "" { + return ctx + } + return context.WithValue(ctx, requestIDContextKey{}, requestID) +} + +func RequestIDFromContext(ctx context.Context) string { + if ctx == nil { + return "" + } + if value, ok := ctx.Value(requestIDContextKey{}).(string); ok { + return NormalizeRequestID(value) + } + return "" +} diff --git a/internal/pkg/tracex/request_id_test.go b/internal/pkg/tracex/request_id_test.go new file mode 100644 index 0000000..0d2f15e --- /dev/null +++ b/internal/pkg/tracex/request_id_test.go @@ -0,0 +1,24 @@ +package tracex + +import "testing" + +func TestNormalizeRequestID(t *testing.T) { + if got := NormalizeRequestID(" trace-123 "); got != "trace-123" { + t.Fatalf("NormalizeRequestID()=%q want %q", got, "trace-123") + } + if got := NormalizeRequestID("bad\nid"); got != "" { + t.Fatalf("NormalizeRequestID()=%q want empty", got) + } + if got := NormalizeRequestID(""); got != "" { + t.Fatalf("NormalizeRequestID()=%q want empty", got) + } +} + +func TestEnsureRequestID(t *testing.T) { + if got := EnsureRequestID("trace-123"); got != "trace-123" { + t.Fatalf("EnsureRequestID(existing)=%q want %q", got, "trace-123") + } + if got := EnsureRequestID(""); got == "" { + t.Fatalf("EnsureRequestID(empty) should generate a value") + } +} diff --git a/internal/services/conversation_event_log_service.go b/internal/services/conversation_event_log_service.go index 45d1876..30be39d 100644 --- a/internal/services/conversation_event_log_service.go +++ b/internal/services/conversation_event_log_service.go @@ -3,6 +3,7 @@ package services import ( "cs-agent/internal/models" "cs-agent/internal/pkg/enums" + "cs-agent/internal/pkg/tracex" "cs-agent/internal/repositories" "strings" "time" @@ -69,8 +70,13 @@ func (s *conversationEventLogService) Delete(id int64) { } func (s *conversationEventLogService) CreateEvent(ctx *sqls.TxContext, conversationID int64, eventType enums.IMEventType, operatorType enums.IMSenderType, operatorID int64, content, payload string) error { + return s.CreateEventWithRequestID(ctx, conversationID, "", eventType, operatorType, operatorID, content, payload) +} + +func (s *conversationEventLogService) CreateEventWithRequestID(ctx *sqls.TxContext, conversationID int64, requestID string, eventType enums.IMEventType, operatorType enums.IMSenderType, operatorID int64, content, payload string) error { return repositories.ConversationEventLogRepository.Create(ctx.Tx, &models.ConversationEventLog{ ConversationID: conversationID, + RequestID: tracex.NormalizeRequestID(requestID), EventType: eventType, OperatorType: operatorType, OperatorID: operatorID, diff --git a/internal/services/conversation_human_dispatch_service.go b/internal/services/conversation_human_dispatch_service.go index 93fb87f..7f2609d 100644 --- a/internal/services/conversation_human_dispatch_service.go +++ b/internal/services/conversation_human_dispatch_service.go @@ -46,6 +46,10 @@ func newConversationHumanDispatchService() *conversationHumanDispatchService { } func (s *conversationHumanDispatchService) TryOffHoursHandoffByAI(conversationID int64, aiAgent models.AIAgent, reason string) (bool, error) { + return s.TryOffHoursHandoffByAIWithRequestID(conversationID, aiAgent, reason, "") +} + +func (s *conversationHumanDispatchService) TryOffHoursHandoffByAIWithRequestID(conversationID int64, aiAgent models.AIAgent, reason string, requestID string) (bool, error) { conversation := ConversationService.Get(conversationID) if conversation == nil { return false, errorsx.InvalidParam("会话不存在") @@ -55,16 +59,20 @@ func (s *conversationHumanDispatchService) TryOffHoursHandoffByAI(conversationID if len(activeTeamIDs) > 0 { return false, nil } - if err := s.createEvent(conversationID, enums.IMEventTypeTransfer, enums.IMSenderTypeAI, aiAgent.ID, "转人工失败:非服务时间", strings.TrimSpace(reason)); err != nil { + if err := s.createEventWithRequestID(conversationID, requestID, enums.IMEventTypeTransfer, enums.IMSenderTypeAI, aiAgent.ID, "转人工失败:非服务时间", strings.TrimSpace(reason)); err != nil { return true, err } - if err := s.sendAIText(conversationID, aiAgent.ID, HandoffOffHoursMessage); err != nil { + if err := s.sendAITextWithRequestID(conversationID, aiAgent.ID, HandoffOffHoursMessage, requestID); err != nil { return true, err } return true, nil } func (s *conversationHumanDispatchService) HandoffByAI(conversationID int64, aiAgent models.AIAgent, reason string) (*HandoffDecisionResult, error) { + return s.HandoffByAIWithRequestID(conversationID, aiAgent, reason, "") +} + +func (s *conversationHumanDispatchService) HandoffByAIWithRequestID(conversationID int64, aiAgent models.AIAgent, reason string, requestID string) (*HandoffDecisionResult, error) { conversation := ConversationService.Get(conversationID) if conversation == nil { return nil, errorsx.InvalidParam("会话不存在") @@ -72,16 +80,16 @@ func (s *conversationHumanDispatchService) HandoffByAI(conversationID int64, aiA teamIDs := orderedPositiveIDs(aiAgent.TeamIDs) activeTeamIDs := ConversationDispatchService.findActiveScheduleTeamIDs(teamIDs, time.Now()) if len(activeTeamIDs) == 0 { - if _, err := s.TryOffHoursHandoffByAI(conversationID, aiAgent, reason); err != nil { + if _, err := s.TryOffHoursHandoffByAIWithRequestID(conversationID, aiAgent, reason, requestID); err != nil { return nil, err } return &HandoffDecisionResult{Decision: HandoffDecisionOffHours, Message: HandoffOffHoursMessage}, nil } - if err := s.markHandoff(conversationID, aiAgent, reason); err != nil { + if err := s.markHandoff(conversationID, aiAgent, reason, requestID); err != nil { return nil, err } - return s.dispatchAfterHandoff(conversationID, aiAgent.ID, activeTeamIDs, strings.TrimSpace(reason), true) + return s.dispatchAfterHandoffWithRequestID(conversationID, aiAgent.ID, activeTeamIDs, strings.TrimSpace(reason), true, requestID) } func (s *conversationHumanDispatchService) ApplyHumanOnlyCreate(conversationID int64, aiAgent models.AIAgent) (*HandoffDecisionResult, error) { @@ -141,7 +149,11 @@ func (s *conversationHumanDispatchService) DispatchPendingConversation(conversat } func (s *conversationHumanDispatchService) dispatchAfterHandoff(conversationID, aiAgentID int64, activeTeamIDs []int64, reason string, publishAssignEvent bool) (*HandoffDecisionResult, error) { - if err := s.sendAIText(conversationID, aiAgentID, HandoffWaitingMessage); err != nil { + return s.dispatchAfterHandoffWithRequestID(conversationID, aiAgentID, activeTeamIDs, reason, publishAssignEvent, "") +} + +func (s *conversationHumanDispatchService) dispatchAfterHandoffWithRequestID(conversationID, aiAgentID int64, activeTeamIDs []int64, reason string, publishAssignEvent bool, requestID string) (*HandoffDecisionResult, error) { + if err := s.sendAITextWithRequestID(conversationID, aiAgentID, HandoffWaitingMessage, requestID); err != nil { return nil, err } @@ -175,7 +187,7 @@ func (s *conversationHumanDispatchService) dispatchAfterHandoff(conversationID, } teamID := activeTeamIDs[0] - teamPoolConversation, err := s.moveToTeamPool(conversationID, teamID, reason) + teamPoolConversation, err := s.moveToTeamPoolWithRequestID(conversationID, teamID, reason, requestID) if err != nil { return nil, err } @@ -185,7 +197,7 @@ func (s *conversationHumanDispatchService) dispatchAfterHandoff(conversationID, return &HandoffDecisionResult{Decision: HandoffDecisionTeamPool, TeamID: teamID, Message: HandoffWaitingMessage}, nil } -func (s *conversationHumanDispatchService) markHandoff(conversationID int64, aiAgent models.AIAgent, reason string) error { +func (s *conversationHumanDispatchService) markHandoff(conversationID int64, aiAgent models.AIAgent, reason string, requestID string) error { now := time.Now() trimmedReason := strings.TrimSpace(reason) return sqls.WithTransaction(func(ctx *sqls.TxContext) error { @@ -201,11 +213,15 @@ func (s *conversationHumanDispatchService) markHandoff(conversationID int64, aiA }); err != nil { return err } - return ConversationEventLogService.CreateEvent(ctx, conversationID, enums.IMEventTypeTransfer, enums.IMSenderTypeAI, aiAgent.ID, "AI转人工", trimmedReason) + return ConversationEventLogService.CreateEventWithRequestID(ctx, conversationID, requestID, enums.IMEventTypeTransfer, enums.IMSenderTypeAI, aiAgent.ID, "AI转人工", trimmedReason) }) } func (s *conversationHumanDispatchService) moveToTeamPool(conversationID, teamID int64, reason string) (*models.Conversation, error) { + return s.moveToTeamPoolWithRequestID(conversationID, teamID, reason, "") +} + +func (s *conversationHumanDispatchService) moveToTeamPoolWithRequestID(conversationID, teamID int64, reason string, requestID string) (*models.Conversation, error) { now := time.Now() var conversation *models.Conversation err := sqls.WithTransaction(func(ctx *sqls.TxContext) error { @@ -226,7 +242,7 @@ func (s *conversationHumanDispatchService) moveToTeamPool(conversationID, teamID }); err != nil { return err } - if err := ConversationEventLogService.CreateEvent(ctx, conversationID, enums.IMEventTypeTransfer, enums.IMSenderTypeSystem, 0, "会话进入客服组待接入", ConversationService.buildEventPayload(map[string]any{ + if err := ConversationEventLogService.CreateEventWithRequestID(ctx, conversationID, requestID, enums.IMEventTypeTransfer, enums.IMSenderTypeSystem, 0, "会话进入客服组待接入", ConversationService.buildEventPayload(map[string]any{ "fromStatus": current.Status, "toStatus": enums.IMConversationStatusPending, "fromAssigneeId": current.CurrentAssigneeID, @@ -278,13 +294,21 @@ func (s *conversationHumanDispatchService) moveToGlobalPool(conversationID int64 } func (s *conversationHumanDispatchService) createEvent(conversationID int64, eventType enums.IMEventType, senderType enums.IMSenderType, senderID int64, content, payload string) error { + return s.createEventWithRequestID(conversationID, "", eventType, senderType, senderID, content, payload) +} + +func (s *conversationHumanDispatchService) createEventWithRequestID(conversationID int64, requestID string, eventType enums.IMEventType, senderType enums.IMSenderType, senderID int64, content, payload string) error { return sqls.WithTransaction(func(ctx *sqls.TxContext) error { - return ConversationEventLogService.CreateEvent(ctx, conversationID, eventType, senderType, senderID, content, payload) + return ConversationEventLogService.CreateEventWithRequestID(ctx, conversationID, requestID, eventType, senderType, senderID, content, payload) }) } func (s *conversationHumanDispatchService) sendAIText(conversationID, aiAgentID int64, content string) error { - _, err := MessageService.SendAIServiceNotice(conversationID, aiAgentID, content) + return s.sendAITextWithRequestID(conversationID, aiAgentID, content, "") +} + +func (s *conversationHumanDispatchService) sendAITextWithRequestID(conversationID, aiAgentID int64, content string, requestID string) error { + _, err := MessageService.SendAIServiceNoticeWithRequestID(conversationID, aiAgentID, content, requestID) return err } diff --git a/internal/services/conversation_service.go b/internal/services/conversation_service.go index 833034f..311d2de 100644 --- a/internal/services/conversation_service.go +++ b/internal/services/conversation_service.go @@ -347,12 +347,17 @@ func (s *conversationService) TransferConversation(conversationID, toUserID int6 } func (s *conversationService) HandoffByAI(conversationID int64, aiAgent models.AIAgent, reason string) error { + return s.HandoffByAIWithRequestID(conversationID, aiAgent, reason, "") +} + +func (s *conversationService) HandoffByAIWithRequestID(conversationID int64, aiAgent models.AIAgent, reason string, requestID string) error { if conversationID <= 0 { return errorsx.InvalidParam("会话不存在") } - _, err := ConversationHumanDispatchService.HandoffByAI(conversationID, aiAgent, reason) + _, err := ConversationHumanDispatchService.HandoffByAIWithRequestID(conversationID, aiAgent, reason, requestID) if err != nil { slog.Warn("schedule-aware ai handoff failed", + "requestId", requestID, "conversation_id", conversationID, "ai_agent_id", aiAgent.ID, "error", err) @@ -361,12 +366,17 @@ func (s *conversationService) HandoffByAI(conversationID int64, aiAgent models.A } func (s *conversationService) TryOffHoursHandoffByAI(conversationID int64, aiAgent models.AIAgent, reason string) (bool, error) { + return s.TryOffHoursHandoffByAIWithRequestID(conversationID, aiAgent, reason, "") +} + +func (s *conversationService) TryOffHoursHandoffByAIWithRequestID(conversationID int64, aiAgent models.AIAgent, reason string, requestID string) (bool, error) { if conversationID <= 0 { return false, errorsx.InvalidParam("会话不存在") } - handled, err := ConversationHumanDispatchService.TryOffHoursHandoffByAI(conversationID, aiAgent, reason) + handled, err := ConversationHumanDispatchService.TryOffHoursHandoffByAIWithRequestID(conversationID, aiAgent, reason, requestID) if err != nil { slog.Warn("off-hours ai handoff failed", + "requestId", requestID, "conversation_id", conversationID, "ai_agent_id", aiAgent.ID, "error", err) diff --git a/internal/services/message_service.go b/internal/services/message_service.go index a598c79..8984fee 100644 --- a/internal/services/message_service.go +++ b/internal/services/message_service.go @@ -6,6 +6,7 @@ import ( "cs-agent/internal/pkg/enums" "cs-agent/internal/pkg/errorsx" "cs-agent/internal/pkg/openidentity" + "cs-agent/internal/pkg/tracex" "cs-agent/internal/pkg/utils" "cs-agent/internal/repositories" "log/slog" @@ -130,18 +131,22 @@ func (s *messageService) GetConversationReadTarget(conversationID, messageID int func (s *messageService) SendMessage(conversationID int64, senderType enums.IMSenderType, reqSenderID int64, clientMsgID string, messageType enums.IMMessageType, content, payload string, operator *dto.AuthPrincipal, external *openidentity.ExternalUser) (*models.Message, error) { switch senderType { case enums.IMSenderTypeAgent: - return s.sendMessage(conversationID, enums.IMSenderTypeAgent, reqSenderID, clientMsgID, messageType, content, payload, operator, nil) + return s.sendMessage(conversationID, enums.IMSenderTypeAgent, reqSenderID, clientMsgID, messageType, content, payload, operator, nil, "") case enums.IMSenderTypeAI: - return s.sendMessage(conversationID, enums.IMSenderTypeAI, reqSenderID, clientMsgID, messageType, content, payload, operator, nil) + return s.sendMessage(conversationID, enums.IMSenderTypeAI, reqSenderID, clientMsgID, messageType, content, payload, operator, nil, "") case enums.IMSenderTypeCustomer: - return s.sendMessage(conversationID, enums.IMSenderTypeCustomer, 0, clientMsgID, messageType, content, payload, nil, external) + return s.sendMessage(conversationID, enums.IMSenderTypeCustomer, 0, clientMsgID, messageType, content, payload, nil, external, "") default: return nil, errorsx.InvalidParam("不支持的发送人类型") } } func (s *messageService) SendAgentMessage(conversationID int64, reqSenderID int64, clientMsgID string, messageType enums.IMMessageType, content, payload string, operator *dto.AuthPrincipal) (*models.Message, error) { - return s.sendMessage(conversationID, enums.IMSenderTypeAgent, reqSenderID, clientMsgID, messageType, content, payload, operator, nil) + return s.SendAgentMessageWithRequestID(conversationID, reqSenderID, clientMsgID, messageType, content, payload, operator, "") +} + +func (s *messageService) SendAgentMessageWithRequestID(conversationID int64, reqSenderID int64, clientMsgID string, messageType enums.IMMessageType, content, payload string, operator *dto.AuthPrincipal, requestID string) (*models.Message, error) { + return s.sendMessage(conversationID, enums.IMSenderTypeAgent, reqSenderID, clientMsgID, messageType, content, payload, operator, nil, requestID) } func (s *messageService) RecallAgentMessage(messageID int64, operator *dto.AuthPrincipal) (*models.Message, error) { @@ -240,10 +245,18 @@ func (s *messageService) RecallAgentMessage(messageID int64, operator *dto.AuthP } func (s *messageService) SendAIMessage(conversationID int64, aiAgentID int64, clientMsgID string, messageType enums.IMMessageType, content, payload string, operator *dto.AuthPrincipal) (*models.Message, error) { - return s.sendMessage(conversationID, enums.IMSenderTypeAI, aiAgentID, clientMsgID, messageType, content, payload, operator, nil) + return s.SendAIMessageWithRequestID(conversationID, aiAgentID, clientMsgID, messageType, content, payload, operator, "") +} + +func (s *messageService) SendAIMessageWithRequestID(conversationID int64, aiAgentID int64, clientMsgID string, messageType enums.IMMessageType, content, payload string, operator *dto.AuthPrincipal, requestID string) (*models.Message, error) { + return s.sendMessage(conversationID, enums.IMSenderTypeAI, aiAgentID, clientMsgID, messageType, content, payload, operator, nil, requestID) } func (s *messageService) SendAIServiceNotice(conversationID int64, aiAgentID int64, content string) (*models.Message, error) { + return s.SendAIServiceNoticeWithRequestID(conversationID, aiAgentID, content, "") +} + +func (s *messageService) SendAIServiceNoticeWithRequestID(conversationID int64, aiAgentID int64, content string, requestID string) (*models.Message, error) { conversation := ConversationService.Get(conversationID) if conversation == nil { return nil, errorsx.InvalidParam("会话不存在") @@ -255,7 +268,7 @@ func (s *messageService) SendAIServiceNotice(conversationID int64, aiAgentID int UserID: 0, Username: "system", Nickname: "system", - }, nil) + }, nil, requestID) } func (s *messageService) createAIWelcomeMessage(ctx *sqls.TxContext, conversation *models.Conversation, aiAgent *models.AIAgent, now time.Time) (*models.Message, error) { @@ -351,12 +364,16 @@ func (s *messageService) createAIWelcomeMessage(ctx *sqls.TxContext, conversatio } func (s *messageService) SendCustomerMessage(conversationID int64, clientMsgID string, messageType enums.IMMessageType, content, payload string, external openidentity.ExternalUser) (*models.Message, error) { + return s.SendCustomerMessageWithRequestID(conversationID, clientMsgID, messageType, content, payload, external, "") +} + +func (s *messageService) SendCustomerMessageWithRequestID(conversationID int64, clientMsgID string, messageType enums.IMMessageType, content, payload string, external openidentity.ExternalUser, requestID string) (*models.Message, error) { ext := external - return s.sendMessage(conversationID, enums.IMSenderTypeCustomer, 0, clientMsgID, messageType, content, payload, nil, &ext) + return s.sendMessage(conversationID, enums.IMSenderTypeCustomer, 0, clientMsgID, messageType, content, payload, nil, &ext, requestID) } func (s *messageService) sendMessage(conversationID int64, senderType enums.IMSenderType, reqSenderID int64, clientMsgID string, - messageType enums.IMMessageType, content, payload string, operator *dto.AuthPrincipal, external *openidentity.ExternalUser) (*models.Message, error) { + messageType enums.IMMessageType, content, payload string, operator *dto.AuthPrincipal, external *openidentity.ExternalUser, requestID string) (*models.Message, error) { if senderType == enums.IMSenderTypeCustomer { if external == nil || strings.TrimSpace(external.ExternalID) == "" { @@ -373,11 +390,11 @@ func (s *messageService) sendMessage(conversationID int64, senderType enums.IMSe if err != nil { return nil, err } - return s.sendValidatedMessage(conversation, senderType, reqSenderID, clientMsgID, messageType, content, payload, operator, external) + return s.sendValidatedMessage(conversation, senderType, reqSenderID, clientMsgID, messageType, content, payload, operator, external, requestID) } func (s *messageService) sendValidatedMessage(conversation *models.Conversation, senderType enums.IMSenderType, reqSenderID int64, clientMsgID string, - messageType enums.IMMessageType, content, payload string, operator *dto.AuthPrincipal, external *openidentity.ExternalUser) (*models.Message, error) { + messageType enums.IMMessageType, content, payload string, operator *dto.AuthPrincipal, external *openidentity.ExternalUser, requestID string) (*models.Message, error) { var err error var summary string @@ -398,6 +415,7 @@ func (s *messageService) sendValidatedMessage(conversation *models.Conversation, var ( now = time.Now() + traceID = tracex.NormalizeRequestID(requestID) auditUserID = int64(0) auditUserName = "" nextSeq = repositories.MessageRepository.NextSeqNo(sqls.DB(), conversation.ID) @@ -412,6 +430,7 @@ func (s *messageService) sendValidatedMessage(conversation *models.Conversation, } message := &models.Message{ ConversationID: conversation.ID, + RequestID: traceID, ClientMsgID: clientMsgID, SenderType: senderType, SenderID: reqSenderID, @@ -487,7 +506,7 @@ func (s *messageService) sendValidatedMessage(conversation *models.Conversation, } // 记录事件日志 - if err := ConversationEventLogService.CreateEvent(ctx, conversation.ID, enums.IMEventTypeMessageSend, senderType, + if err := ConversationEventLogService.CreateEventWithRequestID(ctx, conversation.ID, traceID, enums.IMEventTypeMessageSend, senderType, func() int64 { if operator != nil { return operator.UserID diff --git a/internal/services/message_service_test.go b/internal/services/message_service_test.go index dbfa764..3e86143 100644 --- a/internal/services/message_service_test.go +++ b/internal/services/message_service_test.go @@ -163,6 +163,40 @@ func TestConversationCreateCreatesAIWelcomeMessage(t *testing.T) { } } +func TestSendCustomerMessageStoresRequestIDOnMessageAndEvent(t *testing.T) { + db := setupMessageWelcomeTestDB(t) + aiAgent := createWelcomeTestAIAgent(t, db, "") + external := welcomeTestExternalUser("trace-user") + conversation, err := ConversationService.Create(external, 11, aiAgent.ID) + if err != nil { + t.Fatalf("create conversation: %v", err) + } + + message, err := MessageService.SendCustomerMessageWithRequestID( + conversation.ID, + "client-msg-trace", + enums.IMMessageTypeText, + "hello", + "", + external, + "trace-123", + ) + if err != nil { + t.Fatalf("SendCustomerMessageWithRequestID() error = %v", err) + } + if message.RequestID != "trace-123" { + t.Fatalf("message.RequestID=%q want %q", message.RequestID, "trace-123") + } + + var event models.ConversationEventLog + if err := db.Where("conversation_id = ?", conversation.ID).Order("id DESC").First(&event).Error; err != nil { + t.Fatalf("find event: %v", err) + } + if event.RequestID != "trace-123" { + t.Fatalf("event.RequestID=%q want %q", event.RequestID, "trace-123") + } +} + func TestConversationCreateDoesNotDuplicateWelcomeMessageForExistingConversation(t *testing.T) { db := setupMessageWelcomeTestDB(t) aiAgent := createWelcomeTestAIAgent(t, db, "欢迎咨询") diff --git a/internal/services/ws_realtime_types.go b/internal/services/ws_realtime_types.go index 2370e72..8208b00 100644 --- a/internal/services/ws_realtime_types.go +++ b/internal/services/ws_realtime_types.go @@ -137,6 +137,7 @@ func (e RealtimeResyncRequiredEvent) EventPayload() RealtimeEventPayload { type RealtimeMessageCreatedPayload struct { ConversationID int64 `json:"conversationId,omitempty"` MessageID int64 `json:"messageId,omitempty"` + RequestID string `json:"requestId,omitempty"` Message response.MessageResponse `json:"message,omitempty"` Status enums.IMConversationStatus `json:"status,omitempty"` CurrentAssigneeID int64 `json:"currentAssigneeId,omitempty"` diff --git a/internal/services/ws_service.go b/internal/services/ws_service.go index 30dfb54..1c34b0b 100644 --- a/internal/services/ws_service.go +++ b/internal/services/ws_service.go @@ -299,6 +299,7 @@ func (s *wsService) PublishMessageCreated(conversation *models.Conversation, mes Payload: RealtimeMessageCreatedPayload{ ConversationID: conversation.ID, MessageID: message.ID, + RequestID: message.RequestID, Message: s.buildRealtimeMessage(message), Status: conversation.Status, CurrentAssigneeID: conversation.CurrentAssigneeID, @@ -324,6 +325,7 @@ func (s *wsService) buildRealtimeMessage(item *models.Message) response.MessageR ret := response.MessageResponse{ ID: item.ID, ConversationID: item.ConversationID, + RequestID: item.RequestID, ClientMsgID: item.ClientMsgID, SenderType: item.SenderType, SenderID: item.SenderID, diff --git a/web/components/login-form.tsx b/web/components/login-form.tsx index 2e57820..705f6b1 100644 --- a/web/components/login-form.tsx +++ b/web/components/login-form.tsx @@ -6,7 +6,7 @@ import { startTransition, useEffect, useState } from "react" import { toast } from "sonner" import { useAuth } from "@/components/auth-provider" -import { loginWithPassword } from "@/lib/api/auth" +import { fetchAuthOptions, loginWithPassword, type AuthOptions } from "@/lib/api/auth" import { useI18n } from "@/i18n/provider" import { cn } from "@/lib/utils" import { Button } from "@/components/ui/button" @@ -36,6 +36,10 @@ export function LoginForm({ const { session } = useAuth() const [isPending, setIsPending] = useState(false) const [isWxWorkEnv, setIsWxWorkEnv] = useState(false) + const [authOptions, setAuthOptions] = useState({ + wxworkEnabled: false, + oidcEnabled: false, + }) const nextPath = searchParams.get("next") const wxworkError = searchParams.get("wxworkError") const oidcError = searchParams.get("oidcError") @@ -64,6 +68,26 @@ export function LoginForm({ setIsWxWorkEnv(detectWxWorkEnvironment()) }, []) + useEffect(() => { + let cancelled = false + + void fetchAuthOptions() + .then((options) => { + if (!cancelled) { + setAuthOptions(options) + } + }) + .catch(() => { + if (!cancelled) { + setAuthOptions({ wxworkEnabled: false, oidcEnabled: false }) + } + }) + + return () => { + cancelled = true + } + }, []) + async function handleSubmit(event: React.FormEvent) { event.preventDefault() const formData = new FormData(event.currentTarget) @@ -126,33 +150,37 @@ export function LoginForm({ {isPending ? t("auth.signingIn") : t("auth.signIn")} - - - - - - + {authOptions.wxworkEnabled ? ( + + + + ) : null} + {authOptions.oidcEnabled ? ( + + + + ) : null} ) diff --git a/web/lib/api/auth.ts b/web/lib/api/auth.ts index 65ae6d1..f520bbb 100644 --- a/web/lib/api/auth.ts +++ b/web/lib/api/auth.ts @@ -6,6 +6,17 @@ export type LoginRequest = { password: string } +export type AuthOptions = { + wxworkEnabled: boolean + oidcEnabled: boolean +} + +export async function fetchAuthOptions() { + return request("/api/auth/options", { + skipAuth: true, + }) +} + export async function loginWithPassword(payload: LoginRequest) { const data = await request("/api/auth/login", { method: "POST",