Refactor chat retrieval in Neo API to support filtering and pagination
- Updated GetChats method across conversation implementations to accept a ChatFilter struct, enabling keyword filtering, pagination, and ordering. - Enhanced handleChatList method in the Neo API to construct a filter from query parameters, improving the flexibility of chat retrieval. - Introduced new types (ChatFilter, ChatGroup, ChatGroupResponse) to facilitate structured responses and better organization of chat data. - Improved error handling and response structure for chat retrieval, ensuring robust feedback and clarity in API responses.
This commit is contained in:
parent
f666070bb3
commit
9db67a2456
7 changed files with 245 additions and 25 deletions
26
neo/api.go
26
neo/api.go
|
|
@ -5,6 +5,7 @@ import (
|
||||||
"io"
|
"io"
|
||||||
"net/url"
|
"net/url"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
@ -12,6 +13,7 @@ import (
|
||||||
"github.com/yaoapp/gou/api"
|
"github.com/yaoapp/gou/api"
|
||||||
"github.com/yaoapp/gou/process"
|
"github.com/yaoapp/gou/process"
|
||||||
"github.com/yaoapp/yao/helper"
|
"github.com/yaoapp/yao/helper"
|
||||||
|
"github.com/yaoapp/yao/neo/conversation"
|
||||||
"github.com/yaoapp/yao/neo/message"
|
"github.com/yaoapp/yao/neo/message"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -123,17 +125,33 @@ func (neo *DSL) handleChatList(c *gin.Context) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get keywords from query parameter
|
// Create filter from query parameters
|
||||||
keywords := c.Query("keywords")
|
filter := conversation.ChatFilter{
|
||||||
|
Keywords: c.Query("keywords"),
|
||||||
|
Order: c.Query("order"),
|
||||||
|
}
|
||||||
|
|
||||||
list, err := neo.Conversation.GetChats(sid, keywords)
|
// Parse page and pagesize
|
||||||
|
if page := c.Query("page"); page != "" {
|
||||||
|
if n, err := strconv.Atoi(page); err == nil {
|
||||||
|
filter.Page = n
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if pageSize := c.Query("pagesize"); pageSize != "" {
|
||||||
|
if n, err := strconv.Atoi(pageSize); err == nil {
|
||||||
|
filter.PageSize = n
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
response, err := neo.Conversation.GetChats(sid, filter)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(500, gin.H{"message": err.Error(), "code": 500})
|
c.JSON(500, gin.H{"message": err.Error(), "code": 500})
|
||||||
c.Done()
|
c.Done()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
c.JSON(200, map[string]interface{}{"data": list})
|
c.JSON(200, map[string]interface{}{"data": response})
|
||||||
c.Done()
|
c.Done()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -14,8 +14,14 @@ func (conv *Mongo) UpdateChatTitle(sid string, cid string, title string) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetChats get the chat list
|
// GetChats get the chat list
|
||||||
func (conv *Mongo) GetChats(sid string, keywords ...string) ([]map[string]interface{}, error) {
|
func (conv *Mongo) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, error) {
|
||||||
return []map[string]interface{}{}, nil
|
return &ChatGroupResponse{
|
||||||
|
Groups: []ChatGroup{},
|
||||||
|
Page: filter.Page,
|
||||||
|
PageSize: filter.PageSize,
|
||||||
|
Total: 0,
|
||||||
|
LastPage: 1,
|
||||||
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetHistory get the history
|
// GetHistory get the history
|
||||||
|
|
|
||||||
|
|
@ -14,8 +14,14 @@ func (conv *Redis) UpdateChatTitle(sid string, cid string, title string) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetChats get the chat list
|
// GetChats get the chat list
|
||||||
func (conv *Redis) GetChats(sid string, keywords ...string) ([]map[string]interface{}, error) {
|
func (conv *Redis) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, error) {
|
||||||
return []map[string]interface{}{}, nil
|
return &ChatGroupResponse{
|
||||||
|
Groups: []ChatGroup{},
|
||||||
|
Page: filter.Page,
|
||||||
|
PageSize: filter.PageSize,
|
||||||
|
Total: 0,
|
||||||
|
LastPage: 1,
|
||||||
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetHistory get the history
|
// GetHistory get the history
|
||||||
|
|
|
||||||
|
|
@ -15,10 +15,33 @@ type ChatInfo struct {
|
||||||
History []map[string]interface{} `json:"history"`
|
History []map[string]interface{} `json:"history"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ChatFilter represents the filter parameters for GetChats
|
||||||
|
type ChatFilter struct {
|
||||||
|
Keywords string `json:"keywords,omitempty"`
|
||||||
|
Page int `json:"page,omitempty"` // 页码,从1开始
|
||||||
|
PageSize int `json:"pagesize,omitempty"` // 每页数量
|
||||||
|
Order string `json:"order,omitempty"` // desc/asc
|
||||||
|
}
|
||||||
|
|
||||||
|
// ChatGroup represents a group of chats by date
|
||||||
|
type ChatGroup struct {
|
||||||
|
Label string `json:"label"`
|
||||||
|
Chats []map[string]interface{} `json:"chats"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ChatGroupResponse represents paginated chat groups
|
||||||
|
type ChatGroupResponse struct {
|
||||||
|
Groups []ChatGroup `json:"groups"`
|
||||||
|
Page int `json:"page"` // 当前页码
|
||||||
|
PageSize int `json:"pagesize"` // 每页数量
|
||||||
|
Total int64 `json:"total"` // 总记录数
|
||||||
|
LastPage int `json:"last_page"` // 最后一页页码
|
||||||
|
}
|
||||||
|
|
||||||
// Conversation the store interface
|
// Conversation the store interface
|
||||||
type Conversation interface {
|
type Conversation interface {
|
||||||
UpdateChatTitle(sid string, cid string, title string) error
|
UpdateChatTitle(sid string, cid string, title string) error
|
||||||
GetChats(sid string, keywords ...string) ([]map[string]interface{}, error)
|
GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, error)
|
||||||
GetChat(sid string, cid string) (*ChatInfo, error)
|
GetChat(sid string, cid string) (*ChatInfo, error)
|
||||||
GetHistory(sid string, cid string) ([]map[string]interface{}, error)
|
GetHistory(sid string, cid string) ([]map[string]interface{}, error)
|
||||||
SaveHistory(sid string, messages []map[string]interface{}, cid string) error
|
SaveHistory(sid string, messages []map[string]interface{}, cid string) error
|
||||||
|
|
|
||||||
|
|
@ -14,8 +14,14 @@ func (conv *Weaviate) UpdateChatTitle(sid string, cid string, title string) erro
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetChats get the chat list
|
// GetChats get the chat list
|
||||||
func (conv *Weaviate) GetChats(sid string, keywords ...string) ([]map[string]interface{}, error) {
|
func (conv *Weaviate) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, error) {
|
||||||
return []map[string]interface{}{}, nil
|
return &ChatGroupResponse{
|
||||||
|
Groups: []ChatGroup{},
|
||||||
|
Page: filter.Page,
|
||||||
|
PageSize: filter.PageSize,
|
||||||
|
Total: 0,
|
||||||
|
LastPage: 1,
|
||||||
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetHistory get the history
|
// GetHistory get the history
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ package conversation
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"math"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
|
@ -240,39 +241,124 @@ func (conv *Xun) UpdateChatTitle(sid string, cid string, title string) error {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetChats get the chat list
|
// GetChats get the chat list with grouping by date
|
||||||
func (conv *Xun) GetChats(sid string, keywords ...string) ([]map[string]interface{}, error) {
|
func (conv *Xun) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, error) {
|
||||||
userID, err := conv.getUserID(sid)
|
userID, err := conv.getUserID(sid)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Set defaults
|
||||||
|
if filter.PageSize <= 0 {
|
||||||
|
filter.PageSize = 100
|
||||||
|
}
|
||||||
|
if filter.Page <= 0 {
|
||||||
|
filter.Page = 1
|
||||||
|
}
|
||||||
|
if filter.Order == "" {
|
||||||
|
filter.Order = "desc"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build base query
|
||||||
qb := conv.newQueryChat().
|
qb := conv.newQueryChat().
|
||||||
Select("chat_id", "title").
|
Select("chat_id", "title", "created_at").
|
||||||
Where("sid", userID)
|
Where("sid", userID)
|
||||||
|
|
||||||
// Add title search if keywords provided
|
// Add keyword filter
|
||||||
if len(keywords) > 0 && keywords[0] != "" {
|
if filter.Keywords != "" {
|
||||||
keyword := strings.TrimSpace(keywords[0]) // Trim whitespace from keyword
|
keyword := strings.TrimSpace(filter.Keywords)
|
||||||
if keyword != "" {
|
if keyword != "" {
|
||||||
qb.Where("title", "like", "%"+keyword+"%")
|
qb.Where("title", "like", "%"+keyword+"%")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
rows, err := qb.Get()
|
// Get total count
|
||||||
|
total, err := qb.Clone().Count()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
res := []map[string]interface{}{}
|
// Calculate pagination
|
||||||
for _, row := range rows {
|
offset := (filter.Page - 1) * filter.PageSize
|
||||||
res = append(res, map[string]interface{}{
|
lastPage := int(math.Ceil(float64(total) / float64(filter.PageSize)))
|
||||||
"chat_id": row.Get("chat_id"),
|
|
||||||
"title": row.Get("title"),
|
// Get paginated results
|
||||||
})
|
rows, err := qb.OrderBy("created_at", filter.Order).
|
||||||
|
Offset(offset).
|
||||||
|
Limit(filter.PageSize).
|
||||||
|
Get()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
return res, nil
|
// Group chats by date
|
||||||
|
today := time.Now().Truncate(24 * time.Hour)
|
||||||
|
yesterday := today.AddDate(0, 0, -1)
|
||||||
|
thisWeekStart := today.AddDate(0, 0, -int(today.Weekday()))
|
||||||
|
lastWeekStart := thisWeekStart.AddDate(0, 0, -7)
|
||||||
|
|
||||||
|
groups := map[string][]map[string]interface{}{
|
||||||
|
"Today": {},
|
||||||
|
"Yesterday": {},
|
||||||
|
"This Week": {},
|
||||||
|
"Last Week": {},
|
||||||
|
"Even Earlier": {},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, row := range rows {
|
||||||
|
chat := map[string]interface{}{
|
||||||
|
"chat_id": row.Get("chat_id"),
|
||||||
|
"title": row.Get("title"),
|
||||||
|
}
|
||||||
|
|
||||||
|
createdAt, ok := row.Get("created_at").(time.Time)
|
||||||
|
if !ok {
|
||||||
|
// Try to parse string if it's not already time.Time
|
||||||
|
if timeStr, ok := row.Get("created_at").(string); ok {
|
||||||
|
var err error
|
||||||
|
createdAt, err = time.Parse(time.RFC3339, timeStr)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
createdDate := createdAt.Truncate(24 * time.Hour)
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case createdDate.Equal(today):
|
||||||
|
groups["Today"] = append(groups["Today"], chat)
|
||||||
|
case createdDate.Equal(yesterday):
|
||||||
|
groups["Yesterday"] = append(groups["Yesterday"], chat)
|
||||||
|
case createdDate.After(thisWeekStart) || createdDate.Equal(thisWeekStart):
|
||||||
|
groups["This Week"] = append(groups["This Week"], chat)
|
||||||
|
case createdDate.After(lastWeekStart) || createdDate.Equal(lastWeekStart):
|
||||||
|
groups["Last Week"] = append(groups["Last Week"], chat)
|
||||||
|
default:
|
||||||
|
groups["Even Earlier"] = append(groups["Even Earlier"], chat)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert to ordered slice
|
||||||
|
result := []ChatGroup{}
|
||||||
|
for _, label := range []string{"Today", "Yesterday", "This Week", "Last Week", "Even Earlier"} {
|
||||||
|
if len(groups[label]) > 0 {
|
||||||
|
result = append(result, ChatGroup{
|
||||||
|
Label: label,
|
||||||
|
Chats: groups[label],
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ChatGroupResponse{
|
||||||
|
Groups: result,
|
||||||
|
Page: filter.Page,
|
||||||
|
PageSize: filter.PageSize,
|
||||||
|
Total: total,
|
||||||
|
LastPage: lastPage,
|
||||||
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetHistory get the history
|
// GetHistory get the history
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,9 @@
|
||||||
package conversation
|
package conversation
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/yaoapp/gou/connector"
|
"github.com/yaoapp/gou/connector"
|
||||||
|
|
@ -246,3 +248,76 @@ func TestXunSaveAndGetHistoryWithCID(t *testing.T) {
|
||||||
}
|
}
|
||||||
assert.Equal(t, 2, len(allData))
|
assert.Equal(t, 2, len(allData))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestXunGetChats(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
defer capsule.Schema().DropTableIfExists("__unit_test_conversation")
|
||||||
|
defer capsule.Schema().DropTableIfExists("__unit_test_conversation_chat")
|
||||||
|
|
||||||
|
// Drop both tables before test
|
||||||
|
err := capsule.Schema().DropTableIfExists("__unit_test_conversation")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
err = capsule.Schema().DropTableIfExists("__unit_test_conversation_chat")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
conv, err := NewXun(Setting{
|
||||||
|
Connector: "default",
|
||||||
|
Table: "__unit_test_conversation",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Save some test chats
|
||||||
|
sid := "test_user"
|
||||||
|
messages := []map[string]interface{}{
|
||||||
|
{"role": "user", "content": "test message"},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create chats with different dates
|
||||||
|
for i := 0; i < 5; i++ {
|
||||||
|
chatID := fmt.Sprintf("chat_%d", i)
|
||||||
|
// First create the chat with a title
|
||||||
|
err = conv.newQueryChat().Insert(map[string]interface{}{
|
||||||
|
"chat_id": chatID,
|
||||||
|
"title": fmt.Sprintf("Test Chat %d", i),
|
||||||
|
"sid": sid,
|
||||||
|
"created_at": time.Now(),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Then save the history
|
||||||
|
err = conv.SaveHistory(sid, messages, chatID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test getting chats with default filter
|
||||||
|
filter := ChatFilter{
|
||||||
|
PageSize: 10,
|
||||||
|
Order: "desc",
|
||||||
|
}
|
||||||
|
groups, err := conv.GetChats(sid, filter)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.Greater(t, len(groups.Groups), 0)
|
||||||
|
|
||||||
|
// Test with keywords
|
||||||
|
filter.Keywords = "test"
|
||||||
|
groups, err = conv.GetChats(sid, filter)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.Greater(t, len(groups.Groups), 0)
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue