diff --git a/internal/dao/chat_session.go b/internal/dao/chat_session.go index 0917725d7f..b83cba465b 100644 --- a/internal/dao/chat_session.go +++ b/internal/dao/chat_session.go @@ -221,10 +221,17 @@ func (dao *ChatSessionDAO) ListAgentSessions(ctx context.Context, db *gorm.DB, p if params.Keywords != "" { keywords := strings.ToLower(params.Keywords) escapedKeywords := strings.Trim(strconv.QuoteToASCII(keywords), `"`) + keywordPattern := "%" + keywords + "%" if escapedKeywords == keywords { - query = query.Where("LOWER(message) LIKE ?", "%"+keywords+"%") + query = query.Where("(LOWER(id) LIKE ? OR LOWER(name) LIKE ? OR LOWER(message) LIKE ?)", keywordPattern, keywordPattern, keywordPattern) } else { - query = query.Where("(LOWER(message) LIKE ? OR LOWER(message) LIKE ?)", "%"+keywords+"%", "%"+escapedKeywords+"%") + query = query.Where( + "(LOWER(id) LIKE ? OR LOWER(name) LIKE ? OR LOWER(message) LIKE ? OR LOWER(message) LIKE ?)", + keywordPattern, + keywordPattern, + keywordPattern, + "%"+escapedKeywords+"%", + ) } } diff --git a/internal/dao/chat_session_test.go b/internal/dao/chat_session_test.go index b57eabc05d..19a2ed397c 100644 --- a/internal/dao/chat_session_test.go +++ b/internal/dao/chat_session_test.go @@ -67,6 +67,29 @@ func createAgentSessionForDAOTest(t *testing.T, db *gorm.DB, id, agentID, userID } } +func createNamedAgentSessionForDAOTest(t *testing.T, db *gorm.DB, id, agentID, userID, name string, message json.RawMessage, updateTime int64) { + t.Helper() + + updateDate := time.UnixMilli(updateTime).Local() + session := &entity.API4Conversation{ + ID: id, + Name: &name, + DialogID: agentID, + UserID: userID, + Message: message, + Reference: json.RawMessage(`[]`), + BaseModel: entity.BaseModel{ + CreateTime: &updateTime, + CreateDate: &updateDate, + UpdateTime: &updateTime, + UpdateDate: &updateDate, + }, + } + if err := db.Create(session).Error; err != nil { + t.Fatalf("failed to create session %s: %v", id, err) + } +} + func createChatSessionForDAOTest(t *testing.T, db *gorm.DB, id, chatID, name string, updateTime int64) { t.Helper() @@ -205,3 +228,45 @@ func TestChatSessionDAOListAgentSessionsFiltersAndPaginates(t *testing.T) { t.Fatalf("expected user-1, got %s", sessions[0].UserID) } } + +func TestChatSessionDAOListAgentSessionsSearchesIDNameAndMessage(t *testing.T) { + db := setupChatSessionDAOTestDB(t) + + createNamedAgentSessionForDAOTest(t, db, "release-session-id", "agent-1", "user-1", "plain", json.RawMessage(`[{"content":"ordinary"}]`), 1000) + createNamedAgentSessionForDAOTest(t, db, "session-title", "agent-1", "user-1", "Release Notes", json.RawMessage(`[{"content":"ordinary"}]`), 2000) + createNamedAgentSessionForDAOTest(t, db, "session-message", "agent-1", "user-1", "plain", json.RawMessage(`[{"content":"release details"}]`), 3000) + createNamedAgentSessionForDAOTest(t, db, "other-agent-release", "agent-2", "user-1", "Release Notes", json.RawMessage(`[{"content":"release details"}]`), 4000) + + ctx := t.Context() + total, sessions, err := NewChatSessionDAO().ListAgentSessions(ctx, db, ListAgentSessionsParams{ + AgentID: "agent-1", + Keywords: "release", + Page: 1, + PageSize: 10, + OrderBy: "update_time", + Desc: false, + }) + if err != nil { + t.Fatalf("ListAgentSessions failed: %v", err) + } + + if total != 3 { + t.Fatalf("expected total 3, got %d", total) + } + gotIDs := make([]string, 0, len(sessions)) + for _, session := range sessions { + gotIDs = append(gotIDs, session.ID) + if session.DialogID != "agent-1" { + t.Fatalf("session %s leaked from agent %s", session.ID, session.DialogID) + } + } + wantIDs := []string{"release-session-id", "session-title", "session-message"} + if len(gotIDs) != len(wantIDs) { + t.Fatalf("ids = %v, want %v", gotIDs, wantIDs) + } + for i := range wantIDs { + if gotIDs[i] != wantIDs[i] { + t.Fatalf("ids = %v, want %v", gotIDs, wantIDs) + } + } +}