diff --git a/neo/api.go b/neo/api.go index 5e1b837c..aa2652f8 100644 --- a/neo/api.go +++ b/neo/api.go @@ -28,18 +28,22 @@ func (neo *DSL) API(router *gin.Engine, path string) error { router.OPTIONS(path, neo.optionsHandler) router.OPTIONS(path+"/status", neo.optionsHandler) router.OPTIONS(path+"/chats", neo.optionsHandler) + router.OPTIONS(path+"/chats/:id", neo.optionsHandler) router.OPTIONS(path+"/history", neo.optionsHandler) router.OPTIONS(path+"/upload", neo.optionsHandler) router.OPTIONS(path+"/download", neo.optionsHandler) + router.OPTIONS(path+"/mentions", neo.optionsHandler) // Register endpoints with middlewares router.GET(path, append(middlewares, neo.handleChat)...) router.POST(path, append(middlewares, neo.handleChat)...) router.GET(path+"/status", append(middlewares, neo.handleStatus)...) router.GET(path+"/chats", append(middlewares, neo.handleChatList)...) + router.GET(path+"/chats/:id", append(middlewares, neo.handleChatDetail)...) router.GET(path+"/history", append(middlewares, neo.handleChatHistory)...) router.POST(path+"/upload", append(middlewares, neo.handleUpload)...) router.GET(path+"/download", append(middlewares, neo.handleDownload)...) + router.GET(path+"/mentions", append(middlewares, neo.handleMentions)...) return nil } @@ -107,7 +111,10 @@ func (neo *DSL) handleChatList(c *gin.Context) { return } - list, err := neo.Conversation.GetChats(sid) + // Get keywords from query parameter + keywords := c.Query("keywords") + + list, err := neo.Conversation.GetChats(sid, keywords) if err != nil { c.JSON(500, gin.H{"message": err.Error(), "code": 500}) c.Done() @@ -295,3 +302,52 @@ func (neo *DSL) defaultGuard(c *gin.Context) { c.Set("__sid", user.SID) c.Next() } + +// handleChatDetail handles getting a single chat's details +func (neo *DSL) handleChatDetail(c *gin.Context) { + sid := c.GetString("__sid") + if sid == "" { + c.JSON(400, gin.H{"message": "sid is required", "code": 400}) + c.Done() + return + } + + chatID := c.Param("id") + if chatID == "" { + c.JSON(400, gin.H{"message": "chat id is required", "code": 400}) + c.Done() + return + } + + chat, err := neo.Conversation.GetChat(sid, chatID) + if err != nil { + c.JSON(500, gin.H{"message": err.Error(), "code": 500}) + c.Done() + return + } + + c.JSON(200, chat) + c.Done() +} + +// handleMentions handles getting mentions for a chat +func (neo *DSL) handleMentions(c *gin.Context) { + sid := c.GetString("__sid") + if sid == "" { + c.JSON(400, gin.H{"message": "sid is required", "code": 400}) + c.Done() + return + } + + // Get keywords from query parameter + keywords := c.Query("keywords") + mentions, err := neo.GetMentions(keywords) + if err != nil { + c.JSON(500, gin.H{"message": err.Error(), "code": 500}) + c.Done() + return + } + + c.JSON(200, map[string]interface{}{"data": mentions}) + c.Done() +} diff --git a/neo/conversation/mongo.go b/neo/conversation/mongo.go index 6fe46e50..f2ee634d 100644 --- a/neo/conversation/mongo.go +++ b/neo/conversation/mongo.go @@ -14,7 +14,7 @@ func (conv *Mongo) UpdateChatTitle(sid string, cid string, title string) error { } // GetChats get the chat list -func (conv *Mongo) GetChats(sid string) ([]map[string]interface{}, error) { +func (conv *Mongo) GetChats(sid string, keywords ...string) ([]map[string]interface{}, error) { return []map[string]interface{}{}, nil } @@ -37,3 +37,8 @@ func (conv *Mongo) GetRequest(sid string, rid string) ([]map[string]interface{}, func (conv *Mongo) SaveRequest(sid string, rid string, cid string, messages []map[string]interface{}) error { return nil } + +// GetChat get the chat info and its history +func (conv *Mongo) GetChat(sid string, cid string) (*ChatInfo, error) { + return nil, nil +} diff --git a/neo/conversation/redis.go b/neo/conversation/redis.go index b19412cc..5515dd21 100644 --- a/neo/conversation/redis.go +++ b/neo/conversation/redis.go @@ -14,7 +14,7 @@ func (conv *Redis) UpdateChatTitle(sid string, cid string, title string) error { } // GetChats get the chat list -func (conv *Redis) GetChats(sid string) ([]map[string]interface{}, error) { +func (conv *Redis) GetChats(sid string, keywords ...string) ([]map[string]interface{}, error) { return []map[string]interface{}{}, nil } @@ -37,3 +37,8 @@ func (conv *Redis) GetRequest(sid string, rid string) ([]map[string]interface{}, func (conv *Redis) SaveRequest(sid string, rid string, cid string, messages []map[string]interface{}) error { return nil } + +// GetChat get the chat info and its history +func (conv *Redis) GetChat(sid string, cid string) (*ChatInfo, error) { + return nil, nil +} diff --git a/neo/conversation/types.go b/neo/conversation/types.go index 68ffba17..d28309d9 100644 --- a/neo/conversation/types.go +++ b/neo/conversation/types.go @@ -8,10 +8,17 @@ type Setting struct { TTL int `json:"ttl,omitempty" yaml:"ttl,omitempty"` } +// ChatInfo represents the chat information and its history +type ChatInfo struct { + Chat map[string]interface{} `json:"chat"` + History []map[string]interface{} `json:"history"` +} + // Conversation the store interface type Conversation interface { UpdateChatTitle(sid string, cid string, title string) error - GetChats(sid string) ([]map[string]interface{}, error) + GetChats(sid string, keywords ...string) ([]map[string]interface{}, error) + GetChat(sid string, cid string) (*ChatInfo, error) GetHistory(sid string, cid string) ([]map[string]interface{}, error) SaveHistory(sid string, messages []map[string]interface{}, cid string) error GetRequest(sid string, rid string) ([]map[string]interface{}, error) diff --git a/neo/conversation/weaviate.go b/neo/conversation/weaviate.go index 7495c44d..6770af4c 100644 --- a/neo/conversation/weaviate.go +++ b/neo/conversation/weaviate.go @@ -14,7 +14,7 @@ func (conv *Weaviate) UpdateChatTitle(sid string, cid string, title string) erro } // GetChats get the chat list -func (conv *Weaviate) GetChats(sid string) ([]map[string]interface{}, error) { +func (conv *Weaviate) GetChats(sid string, keywords ...string) ([]map[string]interface{}, error) { return []map[string]interface{}{}, nil } @@ -37,3 +37,8 @@ func (conv *Weaviate) GetRequest(sid string, rid string) ([]map[string]interface func (conv *Weaviate) SaveRequest(sid string, rid string, cid string, messages []map[string]interface{}) error { return nil } + +// GetChat get the chat info and its history +func (conv *Weaviate) GetChat(sid string, cid string) (*ChatInfo, error) { + return nil, nil +} diff --git a/neo/conversation/xun.go b/neo/conversation/xun.go index 0afefd9b..75cd11d6 100644 --- a/neo/conversation/xun.go +++ b/neo/conversation/xun.go @@ -2,6 +2,7 @@ package conversation import ( "fmt" + "strings" "time" "github.com/yaoapp/gou/connector" @@ -29,16 +30,23 @@ type row struct { ExpiredAt interface{} `json:"expired_at"` } +// Public interface methods and constructor remain exported: +// - NewXun +// - UpdateChatTitle +// - GetChats +// - GetChat +// - GetHistory +// - SaveHistory +// - GetRequest +// - SaveRequest + // NewXun create a new conversation func NewXun(setting Setting) (*Xun, error) { - conv := &Xun{setting: setting} if setting.Connector == "default" { conv.query = capsule.Global.Query() conv.schema = capsule.Global.Schema() - } else { - conn, err := connector.Select(setting.Connector) if err != nil { return nil, err @@ -55,7 +63,7 @@ func NewXun(setting Setting) (*Xun, error) { } } - err := conv.Init() + err := conv.initialize() if err != nil { return nil, err } @@ -63,43 +71,175 @@ func NewXun(setting Setting) (*Xun, error) { return conv, nil } -// NewQuery create a new query -func (conv *Xun) NewQuery() query.Query { +// Rename the following functions to start with lowercase letters to make them private: + +func (conv *Xun) newQuery() query.Query { qb := conv.query.New() - qb.Table(conv.setting.Table) + qb.Table(conv.getHistoryTable()) return qb } +func (conv *Xun) newQueryChat() query.Query { + qb := conv.query.New() + qb.Table(conv.getChatTable()) + return qb +} + +func (conv *Xun) clean() { + nums, err := conv.newQuery().Where("expired_at", "<=", time.Now()).Delete() + if err != nil { + log.Error("Clean the conversation table error: %s", err.Error()) + return + } + + if nums > 0 { + log.Trace("Clean the conversation table: %s %d", conv.setting.Table, nums) + } +} + +// Rename Init to initialize to avoid conflicts +func (conv *Xun) initialize() error { + // Initialize history table + if err := conv.initHistoryTable(); err != nil { + return err + } + + // Initialize chat table + if err := conv.initChatTable(); err != nil { + return err + } + + return nil +} + +func (conv *Xun) initHistoryTable() error { + historyTable := conv.getHistoryTable() + has, err := conv.schema.HasTable(historyTable) + if err != nil { + return err + } + + // Create the history table + if !has { + err = conv.schema.CreateTable(historyTable, func(table schema.Blueprint) { + table.ID("id") + table.String("sid", 255).Index() + table.String("rid", 255).Null().Index() + table.String("cid", 200).Null().Index() + table.String("role", 200).Null().Index() + table.String("name", 200).Null().Index() + table.Text("content").Null() + table.TimestampTz("created_at").SetDefaultRaw("NOW()").Index() + table.TimestampTz("updated_at").Null().Index() + table.TimestampTz("expired_at").Null().Index() + }) + + if err != nil { + return err + } + log.Trace("Create the conversation history table: %s", historyTable) + } + + // Validate the table + tab, err := conv.schema.GetTable(historyTable) + if err != nil { + return err + } + + fields := []string{"id", "sid", "rid", "cid", "role", "name", "content", "created_at", "updated_at", "expired_at"} + for _, field := range fields { + if !tab.HasColumn(field) { + return fmt.Errorf("%s is required", field) + } + } + + return nil +} + +func (conv *Xun) initChatTable() error { + chatTable := conv.getChatTable() + has, err := conv.schema.HasTable(chatTable) + if err != nil { + return err + } + + // Create the chat table + if !has { + err = conv.schema.CreateTable(chatTable, func(table schema.Blueprint) { + table.ID("id") + table.String("chat_id", 200).Unique().Index() + table.String("title", 200).Null() + table.String("sid", 255).Index() + table.TimestampTz("created_at").SetDefaultRaw("NOW()").Index() + table.TimestampTz("updated_at").Null().Index() + }) + + if err != nil { + return err + } + log.Trace("Create the chat table: %s", chatTable) + } + + // Validate the table + tab, err := conv.schema.GetTable(chatTable) + if err != nil { + return err + } + + fields := []string{"id", "chat_id", "title", "sid", "created_at", "updated_at"} + for _, field := range fields { + if !tab.HasColumn(field) { + return fmt.Errorf("%s is required", field) + } + } + + return nil +} + +func (conv *Xun) getHistoryTable() string { + return conv.setting.Table +} + +func (conv *Xun) getChatTable() string { + return conv.setting.Table + "_chat" +} + // UpdateChatTitle update the chat title func (conv *Xun) UpdateChatTitle(sid string, cid string, title string) error { - _, err := conv.NewQuery(). - Where("sid", sid).Where("cid", cid). - Update(map[string]interface{}{"title": title}) + _, err := conv.newQueryChat(). + Where("sid", sid). + Where("chat_id", cid). + Update(map[string]interface{}{ + "title": title, + "updated_at": time.Now(), + }) return err } // GetChats get the chat list -func (conv *Xun) GetChats(sid string) ([]map[string]interface{}, error) { - qb := conv.NewQuery(). - Select("cid"). - Where("sid", sid). - GroupBy("cid") +func (conv *Xun) GetChats(sid string, keywords ...string) ([]map[string]interface{}, error) { + qb := conv.newQueryChat(). + Select("chat_id", "title"). + Where("sid", sid) - if conv.setting.TTL > 0 { - qb.Where("expired_at", ">", time.Now()) + // Add title search if keywords provided + if len(keywords) > 0 && keywords[0] != "" { + keyword := strings.TrimSpace(keywords[0]) // Trim whitespace from keyword + if keyword != "" { + qb.Where("title", "like", "%"+keyword+"%") + } } - res := []map[string]interface{}{} - rows, err := qb.Get() if err != nil { return nil, err } + res := []map[string]interface{}{} for _, row := range rows { res = append(res, map[string]interface{}{ - "chat_id": row.Get("cid"), - "title": row.Get("cid"), + "chat_id": row.Get("chat_id"), + "title": row.Get("title"), }) } @@ -109,7 +249,7 @@ func (conv *Xun) GetChats(sid string) ([]map[string]interface{}, error) { // GetHistory get the history func (conv *Xun) GetHistory(sid string, cid string) ([]map[string]interface{}, error) { - qb := conv.NewQuery(). + qb := conv.newQuery(). Select("role", "name", "content"). Where("sid", sid). Where("cid", cid). @@ -143,7 +283,31 @@ func (conv *Xun) GetHistory(sid string, cid string) ([]map[string]interface{}, e // SaveHistory save the history func (conv *Xun) SaveHistory(sid string, messages []map[string]interface{}, cid string) error { + // First ensure chat record exists + exists, err := conv.newQueryChat(). + Where("chat_id", cid). + Where("sid", sid). + Exists() + if err != nil { + return err + } + + if !exists { + // Create new chat record + err = conv.newQueryChat(). + Insert(map[string]interface{}{ + "chat_id": cid, + "sid": sid, + "created_at": time.Now(), + }) + + if err != nil { + return err + } + } + + // Save message history defer conv.clean() var expiredAt interface{} = nil values := []row{} @@ -167,13 +331,13 @@ func (conv *Xun) SaveHistory(sid string, messages []map[string]interface{}, cid values = append(values, value) } - return conv.NewQuery().Insert(values) + return conv.newQuery().Insert(values) } // GetRequest get the request history func (conv *Xun) GetRequest(sid string, rid string) ([]map[string]interface{}, error) { - qb := conv.NewQuery(). + qb := conv.newQuery(). Select("role", "name", "content", "sid"). Where("rid", rid). Where("sid", sid). @@ -232,75 +396,35 @@ func (conv *Xun) SaveRequest(sid string, rid string, cid string, messages []map[ values = append(values, value) } - return conv.NewQuery().Insert(values) + return conv.newQuery().Insert(values) } -func (conv *Xun) clean() { - nums, err := conv.NewQuery().Where("expired_at", "<=", time.Now()).Delete() +// GetChat get the chat info and its history +func (conv *Xun) GetChat(sid string, cid string) (*ChatInfo, error) { + // Get chat info + qb := conv.newQueryChat(). + Select("chat_id", "title"). + Where("sid", sid). + Where("chat_id", cid) + + row, err := qb.First() if err != nil { - log.Error("Clean the conversation table error: %s", err.Error()) - return + return nil, err } - if nums > 0 { - log.Trace("Clean the conversation table: %s %d", conv.setting.Table, nums) + chat := map[string]interface{}{ + "chat_id": row.Get("chat_id"), + "title": row.Get("title"), } -} - -// Init init the conversation -func (conv *Xun) Init() error { - - has, err := conv.schema.HasTable(conv.setting.Table) - if err != nil { - return err - } - - // create the table - if !has { - err = conv.schema.CreateTable(conv.setting.Table, func(table schema.Blueprint) { - - table.ID("id") // The ID field - table.String("sid", 255).Index() // The Session ID - table.String("rid", 255).Null().Index() // The request ID - table.String("cid", 200).Null().Index() // The Chat ID - table.String("role", 200).Null().Index() // The Message role - table.String("name", 200).Null().Index() // The User name - table.String("title", 200).Null().Index() // The Chat title - table.Text("content").Null() - - table.TimestampTz("created_at").SetDefaultRaw("NOW()").Index() - table.TimestampTz("updated_at").Null().Index() - table.TimestampTz("expired_at").Null().Index() - }) - - if err != nil { - return err - } - log.Trace("Create the conversation table: %s", conv.setting.Table) - } - - // validate the table - tab, err := conv.schema.GetTable(conv.setting.Table) - if err != nil { - return err - } - - fields := []string{"id", "sid", "rid", "cid", "role", "name", "content", "created_at", "updated_at", "expired_at"} - for _, field := range fields { - if !tab.HasColumn(field) { - return fmt.Errorf("%s is required", field) - } - } - - // Auto update the title - if !tab.HasColumn("title") { - err = conv.schema.AlterTable(conv.setting.Table, func(table schema.Blueprint) { - table.String("title", 200).Null().Index() - }) - if err != nil { - return err - } - } - - return nil + + // Get chat history + history, err := conv.GetHistory(sid, cid) + if err != nil { + return nil, err + } + + return &ChatInfo{ + Chat: chat, + History: history, + }, nil } diff --git a/neo/hooks.go b/neo/hooks.go index 7066cf05..775f492f 100644 --- a/neo/hooks.go +++ b/neo/hooks.go @@ -194,3 +194,60 @@ func (neo *DSL) HookWrite(ctx Context, messages []map[string]interface{}, respon return result, nil } + +// HookMention query the mention list +func (neo *DSL) HookMention(ctx context.Context, keywords string) ([]Mention, error) { + + // Default Get the assistant list + if neo.MentionHook == "" { + var mentions []Mention + assistants := neo.GetAssistants() + for _, assistant := range assistants { + mentions = append(mentions, Mention{ + ID: assistant.ID, + Name: assistant.Name, + Type: "assistant", + }) + } + + return mentions, nil + } + + // Create a context with 10 second timeout + timeoutCtx, cancel := context.WithTimeout(ctx, 10*time.Second) + defer cancel() + + p, err := process.Of(neo.MentionHook, keywords) + if err != nil { + return nil, err + } + + err = p.WithContext(timeoutCtx).Execute() + if err != nil { + return nil, err + } + defer p.Release() + + // Check if context was canceled + if timeoutCtx.Err() != nil { + return nil, timeoutCtx.Err() + } + + value := p.Value() + if value == nil { + return nil, nil + } + + var list []Mention + bytes, err := jsoniter.Marshal(value) + if err != nil { + return nil, err + } + + err = jsoniter.Unmarshal(bytes, &list) + if err != nil { + return nil, err + } + + return list, nil +} diff --git a/neo/neo.go b/neo/neo.go index 50ca4c3b..affdb80c 100644 --- a/neo/neo.go +++ b/neo/neo.go @@ -1,6 +1,7 @@ package neo import ( + "context" "fmt" "os" "strings" @@ -47,6 +48,16 @@ func (neo *DSL) Answer(ctx Context, question string, c *gin.Context) error { return neo.chat(ast, ctx, messages, c) } +// GetAssistants returns the list of assistants +func (neo *DSL) GetAssistants() []assistant.Assistant { + return neo.AssistantList +} + +// GetMentions returns the mention list +func (neo *DSL) GetMentions(keywords string) ([]Mention, error) { + return neo.HookMention(context.Background(), keywords) +} + // Upload upload a file func (neo *DSL) Upload(ctx Context, c *gin.Context) (*assistant.File, error) { // Get the file diff --git a/neo/types.go b/neo/types.go index 4e7c7c22..0582cba3 100644 --- a/neo/types.go +++ b/neo/types.go @@ -21,6 +21,7 @@ type DSL struct { Create string `json:"create,omitempty" yaml:"create,omitempty"` Write string `json:"write,omitempty" yaml:"write,omitempty"` AssistantListHook string `json:"assistants,omitempty" yaml:"assistants,omitempty"` // Get the assistant list from the hook + MentionHook string `json:"mentions,omitempty"` // Get the mention list from the hook Prompts []assistant.Prompt `json:"prompts,omitempty" yaml:"prompts,omitempty"` Allows []string `json:"allows,omitempty" yaml:"allows,omitempty"` Assistant assistant.API `json:"-" yaml:"-"` // The default assistant @@ -30,6 +31,14 @@ type DSL struct { AssistantMaps map[string]assistant.Assistant `json:"-" yaml:"-"` } +// Mention list +type Mention struct { + ID string `json:"id"` + Name string `json:"name"` + Avatar string `json:"avatar,omitempty"` + Type string `json:"type,omitempty"` +} + // Context the context type Context struct { Sid string `json:"sid" yaml:"-"` // Session ID