Merge pull request #859 from trheyi/main

Add assistant details to chat API responses
This commit is contained in:
Max 2025-02-09 16:54:52 +08:00 committed by GitHub
commit e4375ea33a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 282 additions and 70 deletions

View file

@ -471,18 +471,27 @@ func (neo *DSL) handleChatLatest(c *gin.Context) {
// Create a new chat // Create a new chat
if len(chats.Groups) == 0 || len(chats.Groups[0].Chats) == 0 { if len(chats.Groups) == 0 || len(chats.Groups[0].Chats) == 0 {
ast := neo.Assistant assistantID := neo.Use
assistantID := c.Query("assistant_id") queryAssistantID := c.Query("assistant_id")
if assistantID != "" { if queryAssistantID != "" {
ast, err = assistant.Get(assistantID) assistantID = queryAssistantID
}
// Get the assistant info
ast, err := assistant.Get(assistantID)
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": map[string]interface{}{"placeholder": ast.GetPlaceholder()}}) c.JSON(200, map[string]interface{}{"data": map[string]interface{}{
"placeholder": ast.GetPlaceholder(),
"assistant_id": ast.ID,
"assistant_name": ast.Name,
"assistant_avatar": ast.Avatar,
"assistant_deleteable": neo.Use != ast.ID,
}})
c.Done() c.Done()
return return
} }
@ -502,6 +511,22 @@ func (neo *DSL) handleChatLatest(c *gin.Context) {
return return
} }
// assistant_id is nil return the default assistant
if chat.Chat["assistant_id"] == nil {
chat.Chat["assistant_id"] = neo.Use
// Get the assistant info
ast, err := assistant.Get(neo.Use)
if err != nil {
c.JSON(500, gin.H{"message": err.Error(), "code": 500})
c.Done()
return
}
chat.Chat["assistant_name"] = ast.Name
chat.Chat["assistant_avatar"] = ast.Avatar
}
chat.Chat["assistant_deleteable"] = neo.Use != chat.Chat["assistant_id"]
c.JSON(200, map[string]interface{}{"data": chat}) c.JSON(200, map[string]interface{}{"data": chat})
c.Done() c.Done()
} }
@ -529,6 +554,22 @@ func (neo *DSL) handleChatDetail(c *gin.Context) {
return return
} }
// assistant_id is nil return the default assistant
if chat.Chat["assistant_id"] == nil {
chat.Chat["assistant_id"] = neo.Use
// Get the assistant info
ast, err := assistant.Get(neo.Use)
if err != nil {
c.JSON(500, gin.H{"message": err.Error(), "code": 500})
c.Done()
return
}
chat.Chat["assistant_name"] = ast.Name
chat.Chat["assistant_avatar"] = ast.Avatar
}
chat.Chat["assistant_deleteable"] = neo.Use != chat.Chat["assistant_id"]
c.JSON(200, map[string]interface{}{"data": chat}) c.JSON(200, map[string]interface{}{"data": chat})
c.Done() c.Done()
} }

View file

@ -56,12 +56,12 @@ func (ast *Assistant) Execute(c *gin.Context, ctx chatctx.Context, input string,
} }
// Execute implements the execute functionality // Execute implements the execute functionality
func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input []chatMessage.Message, options map[string]interface{}, contents *chatMessage.Contents) error { func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input []chatMessage.Message, userOptions map[string]interface{}, contents *chatMessage.Contents) error {
if contents == nil { if contents == nil {
contents = chatMessage.NewContents() contents = chatMessage.NewContents()
} }
options = ast.withOptions(options) options := ast.withOptions(userOptions)
// Add RAG and Version support // Add RAG and Version support
ctx.RAG = rag != nil ctx.RAG = rag != nil
@ -78,6 +78,22 @@ func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input []chatM
return err return err
} }
// Update options if provided
if res != nil && res.Options != nil {
options = res.Options
}
// messages
if res != nil && res.Input != nil {
input = res.Input
}
// Handle next action
// It's not used, return the new assistant_id and chat_id
// if res != nil && res.Next != nil {
// return res.Next.Execute(c, ctx, contents)
// }
// Switch to the new assistant if necessary // Switch to the new assistant if necessary
if res != nil && res.AssistantID != ctx.AssistantID { if res != nil && res.AssistantID != ctx.AssistantID {
newAst, err := Get(res.AssistantID) newAst, err := Get(res.AssistantID)
@ -89,22 +105,25 @@ func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input []chatM
Write(c.Writer) Write(c.Writer)
return err return err
} }
*ast = *newAst
// Reset Message Contents
last := input[len(input)-1]
input, err = newAst.withHistory(ctx, last)
if err != nil {
return err
} }
// Handle next action // Reset options
if res != nil && res.Next != nil { options = newAst.withOptions(userOptions)
return res.Next.Execute(c, ctx, contents)
}
// Update options if provided // Update options if provided
if res != nil && res.Options != nil { if res.Options != nil {
options = res.Options options = res.Options
} }
// messages // Update assistant id
if res != nil && res.Input != nil { ctx.AssistantID = res.AssistantID
input = res.Input return newAst.handleChatStream(c, ctx, input, options, contents)
} }
// Only proceed with chat stream if no specific next action was handled // Only proceed with chat stream if no specific next action was handled

View file

@ -185,6 +185,7 @@ func (conv *Xun) initChatTable() error {
table.ID("id") table.ID("id")
table.String("chat_id", 200).Unique().Index() table.String("chat_id", 200).Unique().Index()
table.String("title", 200).Null() table.String("title", 200).Null()
table.String("assistant_id", 200).Null().Index()
table.String("sid", 255).Index() table.String("sid", 255).Index()
table.TimestampTz("created_at").SetDefaultRaw("NOW()").Index() table.TimestampTz("created_at").SetDefaultRaw("NOW()").Index()
table.TimestampTz("updated_at").Null().Index() table.TimestampTz("updated_at").Null().Index()
@ -202,7 +203,7 @@ func (conv *Xun) initChatTable() error {
return err return err
} }
fields := []string{"id", "chat_id", "title", "sid", "created_at", "updated_at"} fields := []string{"id", "chat_id", "title", "assistant_id", "sid", "created_at", "updated_at"}
for _, field := range fields { for _, field := range fields {
if !tab.HasColumn(field) { if !tab.HasColumn(field) {
return fmt.Errorf("%s is required", field) return fmt.Errorf("%s is required", field)
@ -336,7 +337,7 @@ func (conv *Xun) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, er
// Build base query // Build base query
qb := conv.newQueryChat(). qb := conv.newQueryChat().
Select("chat_id", "title", "created_at", "updated_at"). Select("chat_id", "title", "assistant_id", "created_at", "updated_at").
Where("sid", userID). Where("sid", userID).
Where("chat_id", "!=", "") Where("chat_id", "!=", "")
@ -384,6 +385,36 @@ func (conv *Xun) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, er
"Even Earlier": {}, "Even Earlier": {},
} }
// Get assistant details for all chats
assistantIDs := []interface{}{}
assistantMap := make(map[string]map[string]interface{})
for _, row := range rows {
if assistantID := row.Get("assistant_id"); assistantID != nil && assistantID != "" {
assistantIDs = append(assistantIDs, assistantID)
}
}
if len(assistantIDs) > 0 {
assistants, err := conv.query.New().
Table(conv.getAssistantTable()).
Select("assistant_id", "name", "avatar").
WhereIn("assistant_id", assistantIDs).
Get()
if err != nil {
return nil, err
}
for _, assistant := range assistants {
if id := assistant.Get("assistant_id"); id != nil {
assistantMap[fmt.Sprintf("%v", id)] = map[string]interface{}{
"name": assistant.Get("name"),
"avatar": assistant.Get("avatar"),
}
}
}
}
for _, row := range rows { for _, row := range rows {
chatID := row.Get("chat_id") chatID := row.Get("chat_id")
if chatID == nil || chatID == "" { if chatID == nil || chatID == "" {
@ -393,6 +424,15 @@ func (conv *Xun) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, er
chat := map[string]interface{}{ chat := map[string]interface{}{
"chat_id": chatID, "chat_id": chatID,
"title": row.Get("title"), "title": row.Get("title"),
"assistant_id": row.Get("assistant_id"),
}
// Add assistant details if available
if assistantID := row.Get("assistant_id"); assistantID != nil && assistantID != "" {
if assistant, ok := assistantMap[fmt.Sprintf("%v", assistantID)]; ok {
chat["assistant_name"] = assistant["name"]
chat["assistant_avatar"] = assistant["avatar"]
}
} }
var dbDatetime = row.Get("updated_at") var dbDatetime = row.Get("updated_at")
@ -514,6 +554,14 @@ func (conv *Xun) SaveHistory(sid string, messages []map[string]interface{}, cid
return err return err
} }
// Get assistant_id from context
var assistantID interface{} = nil
if context != nil {
if id, ok := context["assistant_id"].(string); ok && id != "" {
assistantID = id
}
}
// First ensure chat record exists // First ensure chat record exists
exists, err := conv.newQueryChat(). exists, err := conv.newQueryChat().
Where("chat_id", cid). Where("chat_id", cid).
@ -530,12 +578,26 @@ func (conv *Xun) SaveHistory(sid string, messages []map[string]interface{}, cid
Insert(map[string]interface{}{ Insert(map[string]interface{}{
"chat_id": cid, "chat_id": cid,
"sid": userID, "sid": userID,
"assistant_id": assistantID,
"created_at": time.Now(), "created_at": time.Now(),
}) })
if err != nil { if err != nil {
return err return err
} }
} else {
// Update assistant_id if it exists
if assistantID != nil {
_, err = conv.newQueryChat().
Where("chat_id", cid).
Where("sid", userID).
Update(map[string]interface{}{
"assistant_id": assistantID,
})
if err != nil {
return err
}
}
} }
// Save message history // Save message history
@ -637,7 +699,7 @@ func (conv *Xun) GetChat(sid string, cid string) (*ChatInfo, error) {
// Get chat info // Get chat info
qb := conv.newQueryChat(). qb := conv.newQueryChat().
Select("chat_id", "title"). Select("chat_id", "title", "assistant_id").
Where("sid", userID). Where("sid", userID).
Where("chat_id", cid) Where("chat_id", cid)
@ -654,6 +716,24 @@ func (conv *Xun) GetChat(sid string, cid string) (*ChatInfo, error) {
chat := map[string]interface{}{ chat := map[string]interface{}{
"chat_id": row.Get("chat_id"), "chat_id": row.Get("chat_id"),
"title": row.Get("title"), "title": row.Get("title"),
"assistant_id": row.Get("assistant_id"),
}
// Get assistant details if assistant_id exists
if assistantID := row.Get("assistant_id"); assistantID != nil && assistantID != "" {
assistant, err := conv.query.New().
Table(conv.getAssistantTable()).
Select("name", "avatar").
Where("assistant_id", assistantID).
First()
if err != nil {
return nil, err
}
if assistant != nil {
chat["assistant_name"] = assistant.Get("name")
chat["assistant_avatar"] = assistant.Get("avatar")
}
} }
// Get chat history // Get chat history
@ -715,24 +795,6 @@ func (conv *Xun) DeleteAllChats(sid string) error {
return err return err
} }
// processJSONField processes a field that should be stored as JSON string
func (conv *Xun) processJSONField(field interface{}) (interface{}, error) {
if field == nil {
return nil, nil
}
switch v := field.(type) {
case string:
return v, nil
default:
jsonStr, err := jsoniter.MarshalToString(v)
if err != nil {
return nil, fmt.Errorf("failed to marshal %v to JSON: %v", field, err)
}
return jsonStr, nil
}
}
// parseJSONFields parses JSON string fields into their corresponding Go types // parseJSONFields parses JSON string fields into their corresponding Go types
func (conv *Xun) parseJSONFields(data map[string]interface{}, fields []string) { func (conv *Xun) parseJSONFields(data map[string]interface{}, fields []string) {
for _, field := range fields { for _, field := range fields {

View file

@ -245,11 +245,15 @@ func TestXunSaveAndGetHistoryWithCID(t *testing.T) {
// save the history with specific cid // save the history with specific cid
sid := "123456" sid := "123456"
cid := "789012" cid := "789012"
assistantID := "test-assistant-1"
messages := []map[string]interface{}{ messages := []map[string]interface{}{
{"role": "user", "name": "user1", "content": "hello"}, {"role": "user", "name": "user1", "content": "hello"},
{"role": "assistant", "name": "assistant1", "content": "Hi! How can I help you?"}, {"role": "assistant", "name": "assistant1", "content": "Hi! How can I help you?"},
} }
err = store.SaveHistory(sid, messages, cid, nil) context := map[string]interface{}{
"assistant_id": assistantID,
}
err = store.SaveHistory(sid, messages, cid, context)
assert.Nil(t, err) assert.Nil(t, err)
// get the history for specific cid // get the history for specific cid
@ -259,14 +263,29 @@ func TestXunSaveAndGetHistoryWithCID(t *testing.T) {
} }
assert.Equal(t, 2, len(data)) assert.Equal(t, 2, len(data))
// save another message with different cid // Verify assistant_id is saved in chat
chat, err := store.GetChat(sid, cid)
assert.Nil(t, err)
assert.Equal(t, assistantID, chat.Chat["assistant_id"])
// save another message with different cid and assistant
anotherCID := "345678" anotherCID := "345678"
anotherAssistantID := "test-assistant-2"
moreMessages := []map[string]interface{}{ moreMessages := []map[string]interface{}{
{"role": "user", "name": "user1", "content": "another message"}, {"role": "user", "name": "user1", "content": "another message"},
{"role": "assistant", "name": "assistant2", "content": "Hello!"},
} }
err = store.SaveHistory(sid, moreMessages, anotherCID, nil) anotherContext := map[string]interface{}{
"assistant_id": anotherAssistantID,
}
err = store.SaveHistory(sid, moreMessages, anotherCID, anotherContext)
assert.Nil(t, err) assert.Nil(t, err)
// Verify second chat's assistant_id
chat2, err := store.GetChat(sid, anotherCID)
assert.Nil(t, err)
assert.Equal(t, anotherAssistantID, chat2.Chat["assistant_id"])
// get history for the first cid - should still be 2 messages // get history for the first cid - should still be 2 messages
data, err = store.GetHistory(sid, cid) data, err = store.GetHistory(sid, cid)
if err != nil { if err != nil {
@ -274,12 +293,12 @@ func TestXunSaveAndGetHistoryWithCID(t *testing.T) {
} }
assert.Equal(t, 2, len(data)) assert.Equal(t, 2, len(data))
// get history for the second cid - should be 1 message // get history for the second cid - should be 2 messages
data, err = store.GetHistory(sid, anotherCID) data, err = store.GetHistory(sid, anotherCID)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
assert.Equal(t, 1, len(data)) assert.Equal(t, 2, len(data))
// get all history for the sid without specifying cid // get all history for the sid without specifying cid
allData, err := store.GetHistory(sid, cid) allData, err := store.GetHistory(sid, cid)
@ -294,8 +313,9 @@ func TestXunGetChats(t *testing.T) {
defer test.Clean() defer test.Clean()
defer capsule.Schema().DropTableIfExists("__unit_test_conversation_history") defer capsule.Schema().DropTableIfExists("__unit_test_conversation_history")
defer capsule.Schema().DropTableIfExists("__unit_test_conversation_chat") defer capsule.Schema().DropTableIfExists("__unit_test_conversation_chat")
defer capsule.Schema().DropTableIfExists("__unit_test_conversation_assistant")
// Drop both tables before test // Drop tables before test
err := capsule.Schema().DropTableIfExists("__unit_test_conversation_history") err := capsule.Schema().DropTableIfExists("__unit_test_conversation_history")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@ -304,6 +324,10 @@ func TestXunGetChats(t *testing.T) {
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
err = capsule.Schema().DropTableIfExists("__unit_test_conversation_assistant")
if err != nil {
t.Fatal(err)
}
store, err := NewXun(Setting{ store, err := NewXun(Setting{
Connector: "default", Connector: "default",
@ -313,45 +337,111 @@ func TestXunGetChats(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
// Create test assistants first
assistant1 := map[string]interface{}{
"assistant_id": "test-assistant-1",
"name": "Test Assistant 1",
"avatar": "avatar1.png",
"type": "assistant",
"connector": "test",
}
assistant2 := map[string]interface{}{
"assistant_id": "test-assistant-2",
"name": "Test Assistant 2",
"avatar": "avatar2.png",
"type": "assistant",
"connector": "test",
}
_, err = store.SaveAssistant(assistant1)
assert.Nil(t, err)
_, err = store.SaveAssistant(assistant2)
assert.Nil(t, err)
// Save some test chats // Save some test chats
sid := "test_user" sid := "test_user"
messages := []map[string]interface{}{ messages := []map[string]interface{}{
{"role": "user", "content": "test message"}, {"role": "user", "content": "test message"},
} }
// Create chats with different dates // Create chats with different dates and assistants
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
chatID := fmt.Sprintf("chat_%d", i) chatID := fmt.Sprintf("chat_%d", i)
title := fmt.Sprintf("Test Chat %d", i) title := fmt.Sprintf("Test Chat %d", i)
var context map[string]interface{}
// Alternate between having assistant and no assistant
if i%2 == 0 {
context = map[string]interface{}{
"assistant_id": "test-assistant-1",
}
} else if i%3 == 0 {
context = map[string]interface{}{
"assistant_id": "test-assistant-2",
}
}
// Save history first to create the chat // Save history first to create the chat
err = store.SaveHistory(sid, messages, chatID, nil) err = store.SaveHistory(sid, messages, chatID, context)
assert.Nil(t, err) assert.Nil(t, err)
// Update the chat title // Update the chat title
err = store.UpdateChatTitle(sid, chatID, title) err = store.UpdateChatTitle(sid, chatID, title)
assert.Nil(t, err) assert.Nil(t, err)
// Verify chat was created with correct assistant info
chat, err := store.GetChat(sid, chatID)
assert.Nil(t, err)
assert.NotNil(t, chat)
assert.Equal(t, chatID, chat.Chat["chat_id"])
assert.Equal(t, title, chat.Chat["title"])
if i%2 == 0 {
assert.Equal(t, "test-assistant-1", chat.Chat["assistant_id"])
assert.Equal(t, "Test Assistant 1", chat.Chat["assistant_name"])
assert.Equal(t, "avatar1.png", chat.Chat["assistant_avatar"])
} else if i%3 == 0 {
assert.Equal(t, "test-assistant-2", chat.Chat["assistant_id"])
assert.Equal(t, "Test Assistant 2", chat.Chat["assistant_name"])
assert.Equal(t, "avatar2.png", chat.Chat["assistant_avatar"])
} else {
assert.Nil(t, chat.Chat["assistant_id"])
assert.Nil(t, chat.Chat["assistant_name"])
assert.Nil(t, chat.Chat["assistant_avatar"])
}
} }
// Test getting chats with default filter // Test GetChats
filter := ChatFilter{ filter := ChatFilter{
PageSize: 10, PageSize: 10,
Order: "desc", Order: "desc",
} }
groups, err := store.GetChats(sid, filter) groups, err := store.GetChats(sid, filter)
if err != nil { assert.Nil(t, err)
t.Fatal(err) assert.NotNil(t, groups)
}
assert.Greater(t, len(groups.Groups), 0) assert.Greater(t, len(groups.Groups), 0)
// Verify assistant information in chat list
for _, group := range groups.Groups {
for _, chat := range group.Chats {
if assistantID, ok := chat["assistant_id"].(string); ok && assistantID != "" {
if assistantID == "test-assistant-1" {
assert.Equal(t, "Test Assistant 1", chat["assistant_name"])
assert.Equal(t, "avatar1.png", chat["assistant_avatar"])
} else if assistantID == "test-assistant-2" {
assert.Equal(t, "Test Assistant 2", chat["assistant_name"])
assert.Equal(t, "avatar2.png", chat["assistant_avatar"])
}
} else {
assert.Nil(t, chat["assistant_name"])
assert.Nil(t, chat["assistant_avatar"])
}
}
}
// Test with keywords // Test with keywords
filter.Keywords = "test" filter.Keywords = "test"
groups, err = store.GetChats(sid, filter) groups, err = store.GetChats(sid, filter)
if err != nil { assert.Nil(t, err)
t.Fatal(err)
}
assert.Greater(t, len(groups.Groups), 0) assert.Greater(t, len(groups.Groups), 0)
} }