Refactor chat storage design to enhance chat and assistant management

- Updated the `CHAT_STORAGE_DESIGN.md` to define the `ChatStore` and `AssistantStore` interfaces, providing clear operations for managing chats, messages, resumes, and assistants.
- Introduced new methods for creating, retrieving, updating, and deleting chats and assistants, along with batch operations for messages and resumes.
- Enhanced the documentation to clarify the responsibilities of each store interface and the associated data structures, ensuring better understanding for future development.
- Revised related functions and tests to support the new design, reinforcing the integrity and performance of chat interactions and assistant management.
This commit is contained in:
Max 2025-12-09 10:44:34 +08:00
parent 81c32ce7a5
commit 3c3177a171
12 changed files with 821 additions and 1236 deletions

View file

@ -693,27 +693,101 @@ func createResumeRecord(ctx *Context, stepType, status string, input, output int
```go ```go
// ChatStore defines the chat storage interface // ChatStore defines the chat storage interface
// Provides operations for chat, message, and resume management
type ChatStore interface { type ChatStore interface {
// ==========================================================================
// Chat Management // Chat Management
// ==========================================================================
// CreateChat creates a new chat session
CreateChat(chat *Chat) error CreateChat(chat *Chat) error
// GetChat retrieves a single chat by ID
GetChat(chatID string) (*Chat, error) GetChat(chatID string) (*Chat, error)
// UpdateChat updates chat fields
UpdateChat(chatID string, updates map[string]interface{}) error UpdateChat(chatID string, updates map[string]interface{}) error
// DeleteChat deletes a chat and its associated messages
DeleteChat(chatID string) error DeleteChat(chatID string) error
// ListChats retrieves a paginated list of chats with optional grouping
ListChats(filter ChatFilter) (*ChatList, error) ListChats(filter ChatFilter) (*ChatList, error)
// ==========================================================================
// Message Management // Message Management
// ==========================================================================
// SaveMessages batch saves messages for a chat
// This is the primary write method - messages are buffered during execution
// and batch-written at the end of a request
SaveMessages(chatID string, messages []*Message) error SaveMessages(chatID string, messages []*Message) error
// GetMessages retrieves messages for a chat with filtering
GetMessages(chatID string, filter MessageFilter) ([]*Message, error) GetMessages(chatID string, filter MessageFilter) ([]*Message, error)
// UpdateMessage updates a single message
UpdateMessage(messageID string, updates map[string]interface{}) error UpdateMessage(messageID string, updates map[string]interface{}) error
// DeleteMessages deletes specific messages from a chat
DeleteMessages(chatID string, messageIDs []string) error DeleteMessages(chatID string, messageIDs []string) error
// ==========================================================================
// Resume Management (only called on failure/interrupt) // Resume Management (only called on failure/interrupt)
// ==========================================================================
// SaveResume batch saves resume records
// Only called when request is interrupted or failed
SaveResume(records []*Resume) error SaveResume(records []*Resume) error
// GetResume retrieves all resume records for a chat
GetResume(chatID string) ([]*Resume, error) GetResume(chatID string) ([]*Resume, error)
// GetLastResume retrieves the last (most recent) resume record for a chat
GetLastResume(chatID string) (*Resume, error) GetLastResume(chatID string) (*Resume, error)
// GetResumeByStackID retrieves resume records for a specific stack
GetResumeByStackID(stackID string) ([]*Resume, error) GetResumeByStackID(stackID string) ([]*Resume, error)
GetStackPath(stackID string) ([]string, error) // Returns [root_stack_id, ..., current_stack_id]
DeleteResume(chatID string) error // Clean up after successful resume // GetStackPath returns the stack path from root to the given stack
// Returns: [root_stack_id, ..., current_stack_id]
GetStackPath(stackID string) ([]string, error)
// DeleteResume deletes all resume records for a chat
// Called after successful resume to clean up
DeleteResume(chatID string) error
}
// AssistantStore defines the assistant storage interface
// Separated from ChatStore for clearer responsibility
type AssistantStore interface {
// SaveAssistant saves assistant information
SaveAssistant(assistant *AssistantModel) (string, error)
// UpdateAssistant updates assistant fields
UpdateAssistant(assistantID string, updates map[string]interface{}) error
// DeleteAssistant deletes an assistant
DeleteAssistant(assistantID string) error
// GetAssistants retrieves a paginated list of assistants with filtering
GetAssistants(filter AssistantFilter, locale ...string) (*AssistantList, error)
// GetAssistantTags retrieves all unique tags from assistants with filtering
GetAssistantTags(filter AssistantFilter, locale ...string) ([]Tag, error)
// GetAssistant retrieves a single assistant by ID
GetAssistant(assistantID string, fields []string, locale ...string) (*AssistantModel, error)
// DeleteAssistants deletes assistants based on filter conditions
DeleteAssistants(filter AssistantFilter) (int64, error)
}
// Store combines ChatStore and AssistantStore interfaces
// This is the main interface for the storage layer
type Store interface {
ChatStore
AssistantStore
} }
// SpaceStore defines the interface for Space snapshot operations // SpaceStore defines the interface for Space snapshot operations
@ -725,7 +799,7 @@ type SpaceStore interface {
// Restore sets multiple key-value pairs from a snapshot // Restore sets multiple key-value pairs from a snapshot
Restore(data map[string]interface{}) error Restore(data map[string]interface{}) error
} }
```` ```
### Data Structures ### Data Structures
@ -736,10 +810,10 @@ type Chat struct {
Title string `json:"title,omitempty"` Title string `json:"title,omitempty"`
AssistantID string `json:"assistant_id"` AssistantID string `json:"assistant_id"`
Mode string `json:"mode"` Mode string `json:"mode"`
Status string `json:"status"` Status string `json:"status"` // "active" or "archived"
Public bool `json:"public"` Public bool `json:"public"` // Whether shared across all teams
Share string `json:"share"` // "private" or "team" Share string `json:"share"` // "private" or "team"
Sort int `json:"sort"` Sort int `json:"sort"` // Sort order for display
LastMessageAt *time.Time `json:"last_message_at,omitempty"` LastMessageAt *time.Time `json:"last_message_at,omitempty"`
Metadata map[string]interface{} `json:"metadata,omitempty"` Metadata map[string]interface{} `json:"metadata,omitempty"`
CreatedAt time.Time `json:"created_at"` CreatedAt time.Time `json:"created_at"`
@ -751,8 +825,8 @@ type Message struct {
MessageID string `json:"message_id"` MessageID string `json:"message_id"`
ChatID string `json:"chat_id"` ChatID string `json:"chat_id"`
RequestID string `json:"request_id,omitempty"` RequestID string `json:"request_id,omitempty"`
Role string `json:"role"` Role string `json:"role"` // "user" or "assistant"
Type string `json:"type"` Type string `json:"type"` // "text", "image", "loading", "tool_call", "retrieval", etc.
Props map[string]interface{} `json:"props"` Props map[string]interface{} `json:"props"`
BlockID string `json:"block_id,omitempty"` BlockID string `json:"block_id,omitempty"`
ThreadID string `json:"thread_id,omitempty"` ThreadID string `json:"thread_id,omitempty"`
@ -763,7 +837,8 @@ type Message struct {
UpdatedAt time.Time `json:"updated_at"` UpdatedAt time.Time `json:"updated_at"`
} }
// Resume represents an execution state for recovery (only stored on failure/interrupt) // Resume represents an execution state for recovery
// Only stored when request is interrupted or failed
type Resume struct { type Resume struct {
ResumeID string `json:"resume_id"` ResumeID string `json:"resume_id"`
ChatID string `json:"chat_id"` ChatID string `json:"chat_id"`
@ -772,7 +847,7 @@ type Resume struct {
StackID string `json:"stack_id"` StackID string `json:"stack_id"`
StackParentID string `json:"stack_parent_id,omitempty"` StackParentID string `json:"stack_parent_id,omitempty"`
StackDepth int `json:"stack_depth"` StackDepth int `json:"stack_depth"`
Type string `json:"type"` Type string `json:"type"` // "input", "hook_create", "llm", "tool", "hook_next", "delegate"
Status string `json:"status"` // "failed" or "interrupted" Status string `json:"status"` // "failed" or "interrupted"
Input map[string]interface{} `json:"input,omitempty"` Input map[string]interface{} `json:"input,omitempty"`
Output map[string]interface{} `json:"output,omitempty"` Output map[string]interface{} `json:"output,omitempty"`
@ -783,6 +858,22 @@ type Resume struct {
CreatedAt time.Time `json:"created_at"` CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"` UpdatedAt time.Time `json:"updated_at"`
} }
// ResumeStatus constants
const (
ResumeStatusFailed = "failed"
ResumeStatusInterrupted = "interrupted"
)
// ResumeType constants
const (
ResumeTypeInput = "input"
ResumeTypeHookCreate = "hook_create"
ResumeTypeLLM = "llm"
ResumeTypeTool = "tool"
ResumeTypeHookNext = "hook_next"
ResumeTypeDelegate = "delegate"
)
``` ```
### Filter Structures ### Filter Structures
@ -811,6 +902,9 @@ type ChatFilter struct {
// Pagination // Pagination
Page int `json:"page,omitempty"` Page int `json:"page,omitempty"`
PageSize int `json:"pagesize,omitempty"` PageSize int `json:"pagesize,omitempty"`
// Permission filter (not serialized)
QueryFilter func(query.Query) `json:"-"` // Custom query function for permission filtering
} }
// MessageFilter for listing messages // MessageFilter for listing messages
@ -1357,3 +1451,4 @@ Main Agent concurrently calls 3 tasks:
- [OpenAPI Request Design](../../openapi/request/REQUEST_DESIGN.md) - Global request tracking, billing, rate limiting - [OpenAPI Request Design](../../openapi/request/REQUEST_DESIGN.md) - Global request tracking, billing, rate limiting
- [Trace Module](../../trace/README.md) - Detailed execution tracing for debugging - [Trace Module](../../trace/README.md) - Detailed execution tracing for debugging
- [Agent Context](../context/README.md) - Context and message handling - [Agent Context](../context/README.md) - Context and message handling
````

View file

@ -10,88 +10,150 @@ func NewMongo() types.Store {
return &Mongo{} return &Mongo{}
} }
// GetChats retrieves a list of chats // =============================================================================
func (m *Mongo) GetChats(sid string, filter types.ChatFilter, locale ...string) (*types.ChatGroupResponse, error) { // Chat Management
return &types.ChatGroupResponse{}, nil // =============================================================================
}
// GetChat retrieves a single chat's information // CreateChat creates a new chat session
func (m *Mongo) GetChat(sid string, cid string, locale ...string) (*types.ChatInfo, error) { func (m *Mongo) CreateChat(chat *types.Chat) error {
return &types.ChatInfo{}, nil // TODO: implement
}
// GetChatWithFilter retrieves a single chat's information with filter options
func (m *Mongo) GetChatWithFilter(sid string, cid string, filter types.ChatFilter, locale ...string) (*types.ChatInfo, error) {
return &types.ChatInfo{}, nil
}
// GetHistory retrieves chat history
func (m *Mongo) GetHistory(sid string, cid string, locale ...string) ([]map[string]interface{}, error) {
return []map[string]interface{}{}, nil
}
// GetHistoryWithFilter retrieves chat history with filter options
func (m *Mongo) GetHistoryWithFilter(sid string, cid string, filter types.ChatFilter, locale ...string) ([]map[string]interface{}, error) {
return []map[string]interface{}{}, nil
}
// SaveHistory saves chat history
func (m *Mongo) SaveHistory(sid string, messages []map[string]interface{}, cid string, context map[string]interface{}) error {
return nil return nil
} }
// DeleteChat deletes a single chat // GetChat retrieves a single chat by ID
func (m *Mongo) DeleteChat(sid string, cid string) error { func (m *Mongo) GetChat(chatID string) (*types.Chat, error) {
// TODO: implement
return nil, nil
}
// UpdateChat updates chat fields
func (m *Mongo) UpdateChat(chatID string, updates map[string]interface{}) error {
// TODO: implement
return nil return nil
} }
// DeleteAllChats deletes all chats // DeleteChat deletes a chat and its associated messages
func (m *Mongo) DeleteAllChats(sid string) error { func (m *Mongo) DeleteChat(chatID string) error {
// TODO: implement
return nil return nil
} }
// UpdateChatTitle updates chat title // ListChats retrieves a paginated list of chats with optional grouping
func (m *Mongo) UpdateChatTitle(sid string, cid string, title string) error { func (m *Mongo) ListChats(filter types.ChatFilter) (*types.ChatList, error) {
// TODO: implement
return nil, nil
}
// =============================================================================
// Message Management
// =============================================================================
// SaveMessages batch saves messages for a chat
func (m *Mongo) SaveMessages(chatID string, messages []*types.Message) error {
// TODO: implement
return nil return nil
} }
// GetMessages retrieves messages for a chat with filtering
func (m *Mongo) GetMessages(chatID string, filter types.MessageFilter) ([]*types.Message, error) {
// TODO: implement
return nil, nil
}
// UpdateMessage updates a single message
func (m *Mongo) UpdateMessage(messageID string, updates map[string]interface{}) error {
// TODO: implement
return nil
}
// DeleteMessages deletes specific messages from a chat
func (m *Mongo) DeleteMessages(chatID string, messageIDs []string) error {
// TODO: implement
return nil
}
// =============================================================================
// Resume Management (only called on failure/interrupt)
// =============================================================================
// SaveResume batch saves resume records
func (m *Mongo) SaveResume(records []*types.Resume) error {
// TODO: implement
return nil
}
// GetResume retrieves all resume records for a chat
func (m *Mongo) GetResume(chatID string) ([]*types.Resume, error) {
// TODO: implement
return nil, nil
}
// GetLastResume retrieves the last resume record for a chat
func (m *Mongo) GetLastResume(chatID string) (*types.Resume, error) {
// TODO: implement
return nil, nil
}
// GetResumeByStackID retrieves resume records for a specific stack
func (m *Mongo) GetResumeByStackID(stackID string) ([]*types.Resume, error) {
// TODO: implement
return nil, nil
}
// GetStackPath returns the stack path from root to the given stack
func (m *Mongo) GetStackPath(stackID string) ([]string, error) {
// TODO: implement
return nil, nil
}
// DeleteResume deletes all resume records for a chat
func (m *Mongo) DeleteResume(chatID string) error {
// TODO: implement
return nil
}
// =============================================================================
// Assistant Management
// =============================================================================
// SaveAssistant saves assistant information // SaveAssistant saves assistant information
func (m *Mongo) SaveAssistant(assistant *types.AssistantModel) (string, error) { func (m *Mongo) SaveAssistant(assistant *types.AssistantModel) (string, error) {
// TODO: implement
return assistant.ID, nil return assistant.ID, nil
} }
// UpdateAssistant updates specific fields of an assistant // UpdateAssistant updates specific fields of an assistant
func (m *Mongo) UpdateAssistant(assistantID string, updates map[string]interface{}) error { func (m *Mongo) UpdateAssistant(assistantID string, updates map[string]interface{}) error {
// TODO: implement
return nil return nil
} }
// DeleteAssistant deletes an assistant // DeleteAssistant deletes an assistant
func (m *Mongo) DeleteAssistant(assistantID string) error { func (m *Mongo) DeleteAssistant(assistantID string) error {
// TODO: implement
return nil return nil
} }
// GetAssistants retrieves a list of assistants // GetAssistants retrieves a list of assistants
func (m *Mongo) GetAssistants(filter types.AssistantFilter, locale ...string) (*types.AssistantList, error) { func (m *Mongo) GetAssistants(filter types.AssistantFilter, locale ...string) (*types.AssistantList, error) {
// TODO: implement
return &types.AssistantList{}, nil return &types.AssistantList{}, nil
} }
// GetAssistant retrieves a single assistant by ID
// fields: Optional list of fields to retrieve. If empty, a default set of fields will be returned.
func (m *Mongo) GetAssistant(assistantID string, fields []string, locale ...string) (*types.AssistantModel, error) {
return nil, nil
}
// DeleteAssistants deletes assistants based on filter conditions (not implemented)
func (m *Mongo) DeleteAssistants(filter types.AssistantFilter) (int64, error) {
return 0, nil
}
// GetAssistantTags retrieves all unique tags from assistants with filtering // GetAssistantTags retrieves all unique tags from assistants with filtering
func (m *Mongo) GetAssistantTags(filter types.AssistantFilter, locale ...string) ([]types.Tag, error) { func (m *Mongo) GetAssistantTags(filter types.AssistantFilter, locale ...string) ([]types.Tag, error) {
// TODO: implement
return []types.Tag{}, nil return []types.Tag{}, nil
} }
// Close closes the store and releases any resources // GetAssistant retrieves a single assistant by ID
func (m *Mongo) Close() error { func (m *Mongo) GetAssistant(assistantID string, fields []string, locale ...string) (*types.AssistantModel, error) {
return nil // TODO: implement
return nil, nil
}
// DeleteAssistants deletes assistants based on filter conditions
func (m *Mongo) DeleteAssistants(filter types.AssistantFilter) (int64, error) {
// TODO: implement
return 0, nil
} }

View file

@ -1,4 +1,4 @@
package store package redis
import "github.com/yaoapp/yao/agent/store/types" import "github.com/yaoapp/yao/agent/store/types"
@ -10,88 +10,150 @@ func NewRedis() types.Store {
return &Redis{} return &Redis{}
} }
// GetChats retrieves a list of chats // =============================================================================
func (r *Redis) GetChats(sid string, filter types.ChatFilter, locale ...string) (*types.ChatGroupResponse, error) { // Chat Management
return &types.ChatGroupResponse{}, nil // =============================================================================
}
// GetChat retrieves a single chat's information // CreateChat creates a new chat session
func (r *Redis) GetChat(sid string, cid string, locale ...string) (*types.ChatInfo, error) { func (r *Redis) CreateChat(chat *types.Chat) error {
return &types.ChatInfo{}, nil // TODO: implement
}
// GetChatWithFilter retrieves a single chat's information with filter options
func (r *Redis) GetChatWithFilter(sid string, cid string, filter types.ChatFilter, locale ...string) (*types.ChatInfo, error) {
return &types.ChatInfo{}, nil
}
// GetHistory retrieves chat history
func (r *Redis) GetHistory(sid string, cid string, locale ...string) ([]map[string]interface{}, error) {
return []map[string]interface{}{}, nil
}
// GetHistoryWithFilter retrieves chat history with filter options
func (r *Redis) GetHistoryWithFilter(sid string, cid string, filter types.ChatFilter, locale ...string) ([]map[string]interface{}, error) {
return []map[string]interface{}{}, nil
}
// SaveHistory saves chat history
func (r *Redis) SaveHistory(sid string, messages []map[string]interface{}, cid string, context map[string]interface{}) error {
return nil return nil
} }
// DeleteChat deletes a single chat // GetChat retrieves a single chat by ID
func (r *Redis) DeleteChat(sid string, cid string) error { func (r *Redis) GetChat(chatID string) (*types.Chat, error) {
// TODO: implement
return nil, nil
}
// UpdateChat updates chat fields
func (r *Redis) UpdateChat(chatID string, updates map[string]interface{}) error {
// TODO: implement
return nil return nil
} }
// DeleteAllChats deletes all chats // DeleteChat deletes a chat and its associated messages
func (r *Redis) DeleteAllChats(sid string) error { func (r *Redis) DeleteChat(chatID string) error {
// TODO: implement
return nil return nil
} }
// UpdateChatTitle updates chat title // ListChats retrieves a paginated list of chats with optional grouping
func (r *Redis) UpdateChatTitle(sid string, cid string, title string) error { func (r *Redis) ListChats(filter types.ChatFilter) (*types.ChatList, error) {
// TODO: implement
return nil, nil
}
// =============================================================================
// Message Management
// =============================================================================
// SaveMessages batch saves messages for a chat
func (r *Redis) SaveMessages(chatID string, messages []*types.Message) error {
// TODO: implement
return nil return nil
} }
// GetMessages retrieves messages for a chat with filtering
func (r *Redis) GetMessages(chatID string, filter types.MessageFilter) ([]*types.Message, error) {
// TODO: implement
return nil, nil
}
// UpdateMessage updates a single message
func (r *Redis) UpdateMessage(messageID string, updates map[string]interface{}) error {
// TODO: implement
return nil
}
// DeleteMessages deletes specific messages from a chat
func (r *Redis) DeleteMessages(chatID string, messageIDs []string) error {
// TODO: implement
return nil
}
// =============================================================================
// Resume Management (only called on failure/interrupt)
// =============================================================================
// SaveResume batch saves resume records
func (r *Redis) SaveResume(records []*types.Resume) error {
// TODO: implement
return nil
}
// GetResume retrieves all resume records for a chat
func (r *Redis) GetResume(chatID string) ([]*types.Resume, error) {
// TODO: implement
return nil, nil
}
// GetLastResume retrieves the last resume record for a chat
func (r *Redis) GetLastResume(chatID string) (*types.Resume, error) {
// TODO: implement
return nil, nil
}
// GetResumeByStackID retrieves resume records for a specific stack
func (r *Redis) GetResumeByStackID(stackID string) ([]*types.Resume, error) {
// TODO: implement
return nil, nil
}
// GetStackPath returns the stack path from root to the given stack
func (r *Redis) GetStackPath(stackID string) ([]string, error) {
// TODO: implement
return nil, nil
}
// DeleteResume deletes all resume records for a chat
func (r *Redis) DeleteResume(chatID string) error {
// TODO: implement
return nil
}
// =============================================================================
// Assistant Management
// =============================================================================
// SaveAssistant saves assistant information // SaveAssistant saves assistant information
func (r *Redis) SaveAssistant(assistant *types.AssistantModel) (string, error) { func (r *Redis) SaveAssistant(assistant *types.AssistantModel) (string, error) {
// TODO: implement
return assistant.ID, nil return assistant.ID, nil
} }
// UpdateAssistant updates specific fields of an assistant // UpdateAssistant updates specific fields of an assistant
func (r *Redis) UpdateAssistant(assistantID string, updates map[string]interface{}) error { func (r *Redis) UpdateAssistant(assistantID string, updates map[string]interface{}) error {
// TODO: implement
return nil return nil
} }
// DeleteAssistant deletes an assistant // DeleteAssistant deletes an assistant
func (r *Redis) DeleteAssistant(assistantID string) error { func (r *Redis) DeleteAssistant(assistantID string) error {
// TODO: implement
return nil return nil
} }
// GetAssistants retrieves a list of assistants // GetAssistants retrieves a list of assistants
func (r *Redis) GetAssistants(filter types.AssistantFilter, locale ...string) (*types.AssistantList, error) { func (r *Redis) GetAssistants(filter types.AssistantFilter, locale ...string) (*types.AssistantList, error) {
// TODO: implement
return &types.AssistantList{}, nil return &types.AssistantList{}, nil
} }
// GetAssistant retrieves a single assistant by ID
// fields: Optional list of fields to retrieve. If empty, a default set of fields will be returned.
func (r *Redis) GetAssistant(assistantID string, fields []string, locale ...string) (*types.AssistantModel, error) {
return nil, nil
}
// DeleteAssistants deletes assistants based on filter conditions (not implemented)
func (r *Redis) DeleteAssistants(filter types.AssistantFilter) (int64, error) {
return 0, nil
}
// GetAssistantTags retrieves all unique tags from assistants with filtering // GetAssistantTags retrieves all unique tags from assistants with filtering
func (r *Redis) GetAssistantTags(filter types.AssistantFilter, locale ...string) ([]types.Tag, error) { func (r *Redis) GetAssistantTags(filter types.AssistantFilter, locale ...string) ([]types.Tag, error) {
// TODO: implement
return []types.Tag{}, nil return []types.Tag{}, nil
} }
// Close closes the store and releases any resources // GetAssistant retrieves a single assistant by ID
func (r *Redis) Close() error { func (r *Redis) GetAssistant(assistantID string, fields []string, locale ...string) (*types.AssistantModel, error) {
return nil // TODO: implement
return nil, nil
}
// DeleteAssistants deletes assistants based on filter conditions
func (r *Redis) DeleteAssistants(filter types.AssistantFilter) (int64, error) {
// TODO: implement
return 0, nil
} }

View file

@ -1,66 +1,108 @@
package types package types
// Store defines the conversation storage interface // ChatStore defines the chat storage interface
// Provides basic operations required for conversation management // Provides operations for chat, message, and resume management
type Store interface { type ChatStore interface {
// GetChats retrieves a list of chats // ==========================================================================
// sid: Session ID // Chat Management
// filter: Filter conditions // ==========================================================================
// Returns: Grouped chat list and potential error
GetChats(sid string, filter ChatFilter, locale ...string) (*ChatGroupResponse, error)
// GetChat retrieves a single chat's information // CreateChat creates a new chat session
// sid: Session ID // chat: Chat session to create
// cid: Chat ID // Returns: Potential error
CreateChat(chat *Chat) error
// GetChat retrieves a single chat by ID
// chatID: Chat ID
// Returns: Chat information and potential error // Returns: Chat information and potential error
GetChat(sid string, cid string, locale ...string) (*ChatInfo, error) GetChat(chatID string) (*Chat, error)
// GetChatWithFilter retrieves a single chat's information with filter options // UpdateChat updates chat fields
// sid: Session ID // chatID: Chat ID
// cid: Chat ID // updates: Map of fields to update
// filter: Filter conditions
// Returns: Chat information and potential error
GetChatWithFilter(sid string, cid string, filter ChatFilter, locale ...string) (*ChatInfo, error)
// GetHistory retrieves chat history
// sid: Session ID
// cid: Chat ID
// Returns: History record list and potential error
GetHistory(sid string, cid string, locale ...string) ([]map[string]interface{}, error)
// GetHistoryWithFilter retrieves chat history with filter options
// sid: Session ID
// cid: Chat ID
// filter: Filter conditions
// Returns: History record list and potential error
GetHistoryWithFilter(sid string, cid string, filter ChatFilter, locale ...string) ([]map[string]interface{}, error)
// SaveHistory saves chat history
// sid: Session ID
// messages: Message list
// cid: Chat ID
// context: Context information
// Returns: Potential error // Returns: Potential error
SaveHistory(sid string, messages []map[string]interface{}, cid string, context map[string]interface{}) error UpdateChat(chatID string, updates map[string]interface{}) error
// DeleteChat deletes a single chat // DeleteChat deletes a chat and its associated messages
// sid: Session ID // chatID: Chat ID
// cid: Chat ID
// Returns: Potential error // Returns: Potential error
DeleteChat(sid string, cid string) error DeleteChat(chatID string) error
// DeleteAllChats deletes all chats // ListChats retrieves a paginated list of chats with optional grouping
// sid: Session ID // filter: Filter conditions including time range, sorting, and grouping
// Returns: Paginated chat list (flat or grouped) and potential error
ListChats(filter ChatFilter) (*ChatList, error)
// ==========================================================================
// Message Management
// ==========================================================================
// SaveMessages batch saves messages for a chat
// This is the primary write method - messages are buffered during execution
// and batch-written at the end of a request
// chatID: Parent chat ID
// messages: Messages to save (includes user input and assistant responses)
// Returns: Potential error // Returns: Potential error
DeleteAllChats(sid string) error SaveMessages(chatID string, messages []*Message) error
// UpdateChatTitle updates chat title // GetMessages retrieves messages for a chat with filtering
// sid: Session ID // chatID: Chat ID
// cid: Chat ID // filter: Filter conditions (role, type, block, thread, etc.)
// title: New title // Returns: Message list and potential error
GetMessages(chatID string, filter MessageFilter) ([]*Message, error)
// UpdateMessage updates a single message
// messageID: Message ID
// updates: Map of fields to update
// Returns: Potential error // Returns: Potential error
UpdateChatTitle(sid string, cid string, title string) error UpdateMessage(messageID string, updates map[string]interface{}) error
// DeleteMessages deletes specific messages from a chat
// chatID: Chat ID
// messageIDs: List of message IDs to delete
// Returns: Potential error
DeleteMessages(chatID string, messageIDs []string) error
// ==========================================================================
// Resume Management (only called on failure/interrupt)
// ==========================================================================
// SaveResume batch saves resume records
// Only called when request is interrupted or failed
// records: Resume records to save
// Returns: Potential error
SaveResume(records []*Resume) error
// GetResume retrieves all resume records for a chat
// chatID: Chat ID
// Returns: Resume records and potential error
GetResume(chatID string) ([]*Resume, error)
// GetLastResume retrieves the last (most recent) resume record for a chat
// chatID: Chat ID
// Returns: Last resume record and potential error
GetLastResume(chatID string) (*Resume, error)
// GetResumeByStackID retrieves resume records for a specific stack
// stackID: Stack ID
// Returns: Resume records and potential error
GetResumeByStackID(stackID string) ([]*Resume, error)
// GetStackPath returns the stack path from root to the given stack
// stackID: Current stack ID
// Returns: Stack path [root_stack_id, ..., current_stack_id] and potential error
GetStackPath(stackID string) ([]string, error)
// DeleteResume deletes all resume records for a chat
// Called after successful resume to clean up
// chatID: Chat ID
// Returns: Potential error
DeleteResume(chatID string) error
}
// AssistantStore defines the assistant storage interface
// Separated from ChatStore for clearer responsibility
type AssistantStore interface {
// SaveAssistant saves assistant information // SaveAssistant saves assistant information
// assistant: Assistant information // assistant: Assistant information
// Returns: Assistant ID and potential error // Returns: Assistant ID and potential error
@ -91,7 +133,7 @@ type Store interface {
// GetAssistant retrieves a single assistant by ID // GetAssistant retrieves a single assistant by ID
// assistantID: Assistant ID // assistantID: Assistant ID
// fields: List of fields to select, empty/nil means default fields (AssistantDefaultFields) // fields: List of fields to select, empty/nil means default fields
// locale: Optional locale for i18n translations // locale: Optional locale for i18n translations
// Returns: Assistant information and potential error // Returns: Assistant information and potential error
GetAssistant(assistantID string, fields []string, locale ...string) (*AssistantModel, error) GetAssistant(assistantID string, fields []string, locale ...string) (*AssistantModel, error)
@ -100,8 +142,21 @@ type Store interface {
// filter: Filter conditions // filter: Filter conditions
// Returns: Number of deleted records and potential error // Returns: Number of deleted records and potential error
DeleteAssistants(filter AssistantFilter) (int64, error) DeleteAssistants(filter AssistantFilter) (int64, error)
}
// Close closes the store and releases any resources
// Returns: Potential error // Store combines ChatStore and AssistantStore interfaces
Close() error // This is the main interface for the storage layer
type Store interface {
ChatStore
AssistantStore
}
// SpaceStore defines the interface for Space snapshot operations
// Note: Space itself uses plan.Space interface, this is for persistence
type SpaceStore interface {
// Snapshot returns all key-value pairs in the space
Snapshot() map[string]interface{}
// Restore sets multiple key-value pairs from a snapshot
Restore(data map[string]interface{}) error
} }

View file

@ -3,6 +3,7 @@ package types
import ( import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"time"
graphragtypes "github.com/yaoapp/gou/graphrag/types" graphragtypes "github.com/yaoapp/gou/graphrag/types"
"github.com/yaoapp/xun/dbal/query" "github.com/yaoapp/xun/dbal/query"
@ -19,40 +20,146 @@ type Setting struct {
Options map[string]interface{} `json:"optional,omitempty" yaml:"optional,omitempty"` // The options for the store Options map[string]interface{} `json:"optional,omitempty" yaml:"optional,omitempty"` // The options for the store
} }
// ChatInfo represents the chat information structure // =============================================================================
// Contains basic information and history for a single chat // Chat Types
type ChatInfo struct { // =============================================================================
Chat map[string]interface{} `json:"chat"` // Basic chat information
History []map[string]interface{} `json:"history"` // Chat history records // Chat represents a chat session
type Chat struct {
ChatID string `json:"chat_id"`
Title string `json:"title,omitempty"`
AssistantID string `json:"assistant_id"`
Mode string `json:"mode"`
Status string `json:"status"` // "active" or "archived"
Public bool `json:"public"` // Whether shared across all teams
Share string `json:"share"` // "private" or "team"
Sort int `json:"sort"` // Sort order for display
LastMessageAt *time.Time `json:"last_message_at,omitempty"`
Metadata map[string]interface{} `json:"metadata,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
} }
// ChatFilter represents the chat filter structure // ChatFilter for listing chats
// Used for filtering and pagination when retrieving chat lists
type ChatFilter struct { type ChatFilter struct {
Keywords string `json:"keywords,omitempty"` // Keyword search UserID string `json:"user_id,omitempty"`
Page int `json:"page,omitempty"` // Page number, starting from 1 TeamID string `json:"team_id,omitempty"`
PageSize int `json:"pagesize,omitempty"` // Number of items per page AssistantID string `json:"assistant_id,omitempty"`
Order string `json:"order,omitempty"` // Sort order: desc/asc Status string `json:"status,omitempty"`
Silent *bool `json:"silent,omitempty"` // Include silent messages (default: false) Keywords string `json:"keywords,omitempty"`
// Time range filter
StartTime *time.Time `json:"start_time,omitempty"` // Filter chats after this time
EndTime *time.Time `json:"end_time,omitempty"` // Filter chats before this time
TimeField string `json:"time_field,omitempty"` // Field for time filter: "created_at" or "last_message_at" (default)
// Sorting
OrderBy string `json:"order_by,omitempty"` // Field to sort by (default: "last_message_at")
Order string `json:"order,omitempty"` // Sort order: "desc" (default) or "asc"
// Response format
GroupBy string `json:"group_by,omitempty"` // "time" for time-based groups, empty for flat list
// Pagination
Page int `json:"page,omitempty"`
PageSize int `json:"pagesize,omitempty"`
// Permission filter (not serialized)
QueryFilter func(query.Query) `json:"-"` // Custom query function for permission filtering
} }
// ChatGroup represents the chat group structure // ChatList paginated response with time-based grouping
// Groups chats by date type ChatList struct {
Data []*Chat `json:"data"`
Groups []*ChatGroup `json:"groups,omitempty"` // Time-based groups for UI display
Page int `json:"page"`
PageSize int `json:"pagesize"`
PageCount int `json:"pagecount"`
Total int `json:"total"`
}
// ChatGroup represents a time-based group of chats
type ChatGroup struct { type ChatGroup struct {
Label string `json:"label"` // Group label (typically a date) Label string `json:"label"` // "Today", "Yesterday", "This Week", "This Month", "Earlier"
Chats []map[string]interface{} `json:"chats"` // List of chats in this group Key string `json:"key"` // "today", "yesterday", "this_week", "this_month", "earlier"
Chats []*Chat `json:"chats"` // Chats in this group
Count int `json:"count"` // Number of chats in group
} }
// ChatGroupResponse represents the paginated chat group response // =============================================================================
// Contains paginated chat group information // Message Types
type ChatGroupResponse struct { // =============================================================================
Groups []ChatGroup `json:"groups"` // List of chat groups
Page int `json:"page"` // Current page number // Message represents a chat message
PageSize int `json:"pagesize"` // Items per page type Message struct {
Total int64 `json:"total"` // Total number of records MessageID string `json:"message_id"`
LastPage int `json:"last_page"` // Last page number ChatID string `json:"chat_id"`
RequestID string `json:"request_id,omitempty"`
Role string `json:"role"` // "user" or "assistant"
Type string `json:"type"` // "text", "image", "loading", "tool_call", "retrieval", etc.
Props map[string]interface{} `json:"props"`
BlockID string `json:"block_id,omitempty"`
ThreadID string `json:"thread_id,omitempty"`
AssistantID string `json:"assistant_id,omitempty"`
Sequence int `json:"sequence"`
Metadata map[string]interface{} `json:"metadata,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
} }
// MessageFilter for listing messages
type MessageFilter struct {
RequestID string `json:"request_id,omitempty"`
Role string `json:"role,omitempty"`
BlockID string `json:"block_id,omitempty"`
ThreadID string `json:"thread_id,omitempty"`
Type string `json:"type,omitempty"`
Limit int `json:"limit,omitempty"`
Offset int `json:"offset,omitempty"`
}
// =============================================================================
// Resume Types (for recovery from interruption/failure)
// =============================================================================
// Resume represents an execution state for recovery
// Only stored when request is interrupted or failed
type Resume struct {
ResumeID string `json:"resume_id"`
ChatID string `json:"chat_id"`
RequestID string `json:"request_id"`
AssistantID string `json:"assistant_id"`
StackID string `json:"stack_id"`
StackParentID string `json:"stack_parent_id,omitempty"`
StackDepth int `json:"stack_depth"`
Type string `json:"type"` // "input", "hook_create", "llm", "tool", "hook_next", "delegate"
Status string `json:"status"` // "failed" or "interrupted"
Input map[string]interface{} `json:"input,omitempty"`
Output map[string]interface{} `json:"output,omitempty"`
SpaceSnapshot map[string]interface{} `json:"space_snapshot,omitempty"` // Shared space data for recovery
Error string `json:"error,omitempty"`
Sequence int `json:"sequence"`
Metadata map[string]interface{} `json:"metadata,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// ResumeStatus constants
const (
ResumeStatusFailed = "failed"
ResumeStatusInterrupted = "interrupted"
)
// ResumeType constants
const (
ResumeTypeInput = "input"
ResumeTypeHookCreate = "hook_create"
ResumeTypeLLM = "llm"
ResumeTypeTool = "tool"
ResumeTypeHookNext = "hook_next"
ResumeTypeDelegate = "delegate"
)
// AssistantFilter represents the assistant filter structure // AssistantFilter represents the assistant filter structure
// Used for filtering and pagination when retrieving assistant lists // Used for filtering and pagination when retrieving assistant lists
type AssistantFilter struct { type AssistantFilter struct {

View file

@ -14,7 +14,7 @@ import (
) )
// SaveAssistant saves assistant information // SaveAssistant saves assistant information
func (conv *Xun) SaveAssistant(assistant *types.AssistantModel) (string, error) { func (store *Xun) SaveAssistant(assistant *types.AssistantModel) (string, error) {
if assistant == nil { if assistant == nil {
return "", fmt.Errorf("assistant cannot be nil") return "", fmt.Errorf("assistant cannot be nil")
} }
@ -33,15 +33,15 @@ func (conv *Xun) SaveAssistant(assistant *types.AssistantModel) (string, error)
// Generate assistant_id if not provided // Generate assistant_id if not provided
if assistant.ID == "" { if assistant.ID == "" {
var err error var err error
assistant.ID, err = conv.GenerateAssistantID() assistant.ID, err = store.GenerateAssistantID()
if err != nil { if err != nil {
return "", err return "", err
} }
} }
// Check if assistant exists // Check if assistant exists
exists, err := conv.query.New(). exists, err := store.query.New().
Table(conv.getAssistantTable()). Table(store.getAssistantTable()).
Where("assistant_id", assistant.ID). Where("assistant_id", assistant.ID).
Exists() Exists()
if err != nil { if err != nil {
@ -197,8 +197,8 @@ func (conv *Xun) SaveAssistant(assistant *types.AssistantModel) (string, error)
// Update or insert // Update or insert
if exists { if exists {
_, err := conv.query.New(). _, err := store.query.New().
Table(conv.getAssistantTable()). Table(store.getAssistantTable()).
Where("assistant_id", assistant.ID). Where("assistant_id", assistant.ID).
Update(data) Update(data)
if err != nil { if err != nil {
@ -207,8 +207,8 @@ func (conv *Xun) SaveAssistant(assistant *types.AssistantModel) (string, error)
return assistant.ID, nil return assistant.ID, nil
} }
err = conv.query.New(). err = store.query.New().
Table(conv.getAssistantTable()). Table(store.getAssistantTable()).
Insert(data) Insert(data)
if err != nil { if err != nil {
return "", err return "", err
@ -217,7 +217,7 @@ func (conv *Xun) SaveAssistant(assistant *types.AssistantModel) (string, error)
} }
// UpdateAssistant updates specific fields of an assistant // UpdateAssistant updates specific fields of an assistant
func (conv *Xun) UpdateAssistant(assistantID string, updates map[string]interface{}) error { func (store *Xun) UpdateAssistant(assistantID string, updates map[string]interface{}) error {
if assistantID == "" { if assistantID == "" {
return fmt.Errorf("assistant_id is required") return fmt.Errorf("assistant_id is required")
} }
@ -226,8 +226,8 @@ func (conv *Xun) UpdateAssistant(assistantID string, updates map[string]interfac
} }
// Check if assistant exists // Check if assistant exists
exists, err := conv.query.New(). exists, err := store.query.New().
Table(conv.getAssistantTable()). Table(store.getAssistantTable()).
Where("assistant_id", assistantID). Where("assistant_id", assistantID).
Exists() Exists()
if err != nil { if err != nil {
@ -291,8 +291,8 @@ func (conv *Xun) UpdateAssistant(assistantID string, updates map[string]interfac
} }
// Perform update // Perform update
_, err = conv.query.New(). _, err = store.query.New().
Table(conv.getAssistantTable()). Table(store.getAssistantTable()).
Where("assistant_id", assistantID). Where("assistant_id", assistantID).
Update(data) Update(data)
@ -300,10 +300,10 @@ func (conv *Xun) UpdateAssistant(assistantID string, updates map[string]interfac
} }
// DeleteAssistant deletes an assistant by assistant_id // DeleteAssistant deletes an assistant by assistant_id
func (conv *Xun) DeleteAssistant(assistantID string) error { func (store *Xun) DeleteAssistant(assistantID string) error {
// Check if assistant exists // Check if assistant exists
exists, err := conv.query.New(). exists, err := store.query.New().
Table(conv.getAssistantTable()). Table(store.getAssistantTable()).
Where("assistant_id", assistantID). Where("assistant_id", assistantID).
Exists() Exists()
if err != nil { if err != nil {
@ -314,17 +314,17 @@ func (conv *Xun) DeleteAssistant(assistantID string) error {
return fmt.Errorf("assistant %s not found", assistantID) return fmt.Errorf("assistant %s not found", assistantID)
} }
_, err = conv.query.New(). _, err = store.query.New().
Table(conv.getAssistantTable()). Table(store.getAssistantTable()).
Where("assistant_id", assistantID). Where("assistant_id", assistantID).
Delete() Delete()
return err return err
} }
// GetAssistants retrieves assistants with pagination and filtering // GetAssistants retrieves assistants with pagination and filtering
func (conv *Xun) GetAssistants(filter types.AssistantFilter, locale ...string) (*types.AssistantList, error) { func (store *Xun) GetAssistants(filter types.AssistantFilter, locale ...string) (*types.AssistantList, error) {
qb := conv.query.New(). qb := store.query.New().
Table(conv.getAssistantTable()) Table(store.getAssistantTable())
// Apply tag filter if provided // Apply tag filter if provided
if len(filter.Tags) > 0 { if len(filter.Tags) > 0 {
@ -450,7 +450,7 @@ func (conv *Xun) GetAssistants(filter types.AssistantFilter, locale ...string) (
} }
// Parse JSON fields // Parse JSON fields
conv.parseJSONFields(data, jsonFields) store.parseJSONFields(data, jsonFields)
// Convert map to types.AssistantModel using existing helper function // Convert map to types.AssistantModel using existing helper function
model, err := types.ToAssistantModel(data) model, err := types.ToAssistantModel(data)
@ -461,7 +461,7 @@ func (conv *Xun) GetAssistants(filter types.AssistantFilter, locale ...string) (
// Apply i18n translations if locale is provided // Apply i18n translations if locale is provided
if len(locale) > 0 && locale[0] != "" && model != nil { if len(locale) > 0 && locale[0] != "" && model != nil {
conv.translate(model, model.ID, locale[0]) store.translate(model, model.ID, locale[0])
} }
assistants = append(assistants, model) assistants = append(assistants, model)
@ -479,9 +479,9 @@ func (conv *Xun) GetAssistants(filter types.AssistantFilter, locale ...string) (
} }
// GetAssistant retrieves a single assistant by ID // GetAssistant retrieves a single assistant by ID
func (conv *Xun) GetAssistant(assistantID string, fields []string, locale ...string) (*types.AssistantModel, error) { func (store *Xun) GetAssistant(assistantID string, fields []string, locale ...string) (*types.AssistantModel, error) {
qb := conv.query.New(). qb := store.query.New().
Table(conv.getAssistantTable()). Table(store.getAssistantTable()).
Where("assistant_id", assistantID) Where("assistant_id", assistantID)
// Apply select fields with security validation // Apply select fields with security validation
@ -515,7 +515,7 @@ func (conv *Xun) GetAssistant(assistantID string, fields []string, locale ...str
// Parse JSON fields // Parse JSON fields
jsonFields := []string{"tags", "modes", "options", "prompts", "prompt_presets", "connector_options", "workflow", "kb", "db", "mcp", "placeholder", "locales", "uses"} jsonFields := []string{"tags", "modes", "options", "prompts", "prompt_presets", "connector_options", "workflow", "kb", "db", "mcp", "placeholder", "locales", "uses"}
conv.parseJSONFields(data, jsonFields) store.parseJSONFields(data, jsonFields)
// Convert map to types.AssistantModel // Convert map to types.AssistantModel
model := &types.AssistantModel{ model := &types.AssistantModel{
@ -661,16 +661,16 @@ func (conv *Xun) GetAssistant(assistantID string, fields []string, locale ...str
// Apply i18n translation if locale is provided // Apply i18n translation if locale is provided
if len(locale) > 0 && locale[0] != "" { if len(locale) > 0 && locale[0] != "" {
conv.translate(model, assistantID, locale[0]) store.translate(model, assistantID, locale[0])
} }
return model, nil return model, nil
} }
// DeleteAssistants deletes assistants based on filter conditions // DeleteAssistants deletes assistants based on filter conditions
func (conv *Xun) DeleteAssistants(filter types.AssistantFilter) (int64, error) { func (store *Xun) DeleteAssistants(filter types.AssistantFilter) (int64, error) {
qb := conv.query.New(). qb := store.query.New().
Table(conv.getAssistantTable()) Table(store.getAssistantTable())
// Apply tag filter if provided // Apply tag filter if provided
if len(filter.Tags) > 0 { if len(filter.Tags) > 0 {
@ -729,8 +729,8 @@ func (conv *Xun) DeleteAssistants(filter types.AssistantFilter) (int64, error) {
} }
// GetAssistantTags retrieves all unique tags from assistants with filtering // GetAssistantTags retrieves all unique tags from assistants with filtering
func (conv *Xun) GetAssistantTags(filter types.AssistantFilter, locale ...string) ([]types.Tag, error) { func (store *Xun) GetAssistantTags(filter types.AssistantFilter, locale ...string) ([]types.Tag, error) {
qb := conv.query.New().Table(conv.getAssistantTable()) qb := store.query.New().Table(store.getAssistantTable())
// Apply type filter (default to "assistant") // Apply type filter (default to "assistant")
typeFilter := "assistant" typeFilter := "assistant"
@ -803,7 +803,7 @@ func (conv *Xun) GetAssistantTags(filter types.AssistantFilter, locale ...string
} }
// translate applies i18n translation to assistant model fields // translate applies i18n translation to assistant model fields
func (conv *Xun) translate(model *types.AssistantModel, assistantID string, locale string) { func (store *Xun) translate(model *types.AssistantModel, assistantID string, locale string) {
if model == nil { if model == nil {
return return
} }

View file

@ -1,4 +1,4 @@
package xun package xun_test
import ( import (
"fmt" "fmt"
@ -11,6 +11,7 @@ import (
"github.com/yaoapp/yao/agent/context" "github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/agent/i18n" "github.com/yaoapp/yao/agent/i18n"
"github.com/yaoapp/yao/agent/store/types" "github.com/yaoapp/yao/agent/store/types"
"github.com/yaoapp/yao/agent/store/xun"
"github.com/yaoapp/yao/config" "github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/test" "github.com/yaoapp/yao/test"
) )
@ -31,13 +32,12 @@ func TestSaveAssistant(t *testing.T) {
defer test.Clean() defer test.Clean()
// Create a new xun store // Create a new xun store
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("CreateNewAssistant", func(t *testing.T) { t.Run("CreateNewAssistant", func(t *testing.T) {
assistant := &types.AssistantModel{ assistant := &types.AssistantModel{
@ -648,13 +648,12 @@ func TestDeleteAssistant(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("DeleteExistingAssistant", func(t *testing.T) { t.Run("DeleteExistingAssistant", func(t *testing.T) {
// Create assistant // Create assistant
@ -696,13 +695,12 @@ func TestGetAssistant(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("GetExistingAssistant", func(t *testing.T) { t.Run("GetExistingAssistant", func(t *testing.T) {
// Create assistant // Create assistant
@ -764,13 +762,12 @@ func TestGetAssistants(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
// Clean up existing data before creating test assistants // Clean up existing data before creating test assistants
deleted, err := store.DeleteAssistants(types.AssistantFilter{}) deleted, err := store.DeleteAssistants(types.AssistantFilter{})
@ -1078,13 +1075,12 @@ func TestDeleteAssistants(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("DeleteByTag", func(t *testing.T) { t.Run("DeleteByTag", func(t *testing.T) {
// Create assistants with specific tag // Create assistants with specific tag
@ -1205,13 +1201,12 @@ func TestGetAssistantTags(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("GetUniqueTags", func(t *testing.T) { t.Run("GetUniqueTags", func(t *testing.T) {
// Create assistants with various tags // Create assistants with various tags
@ -1433,57 +1428,17 @@ func TestGetAssistantTags(t *testing.T) {
}) })
} }
// TestGenerateAssistantID tests the ID generation function
func TestGenerateAssistantID(t *testing.T) {
test.Prepare(t, config.Conf)
defer test.Clean()
store, err := NewXun(types.Setting{
Connector: "default",
})
if err != nil {
t.Fatalf("Failed to create store: %v", err)
}
defer store.Close()
xunStore := store.(*Xun)
t.Run("GenerateUniqueIDs", func(t *testing.T) {
ids := make(map[string]bool)
for i := 0; i < 10; i++ {
id, err := xunStore.GenerateAssistantID()
if err != nil {
t.Fatalf("Failed to generate ID: %v", err)
}
// Verify ID format (6 digits)
if len(id) != 6 {
t.Errorf("Expected 6-digit ID, got %s (length %d)", id, len(id))
}
// Verify ID is unique
if ids[id] {
t.Errorf("Generated duplicate ID: %s", id)
}
ids[id] = true
}
t.Logf("Generated %d unique IDs", len(ids))
})
}
// TestAssistantPermissionFields tests permission management fields // TestAssistantPermissionFields tests permission management fields
func TestAssistantPermissionFields(t *testing.T) { func TestAssistantPermissionFields(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("SaveWithPermissionFields", func(t *testing.T) { t.Run("SaveWithPermissionFields", func(t *testing.T) {
assistant := &types.AssistantModel{ assistant := &types.AssistantModel{
@ -1608,13 +1563,12 @@ func TestEmptyStringAsNull(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("EmptyStringsStoredAsNull", func(t *testing.T) { t.Run("EmptyStringsStoredAsNull", func(t *testing.T) {
// Create assistant with empty strings for nullable fields // Create assistant with empty strings for nullable fields
@ -1711,13 +1665,12 @@ func TestGetAssistantWithLocale(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("GetAssistantWithLocaleTranslation", func(t *testing.T) { t.Run("GetAssistantWithLocaleTranslation", func(t *testing.T) {
// Create assistant with i18n locales // Create assistant with i18n locales
@ -1837,13 +1790,12 @@ func TestGetAssistantsWithLocale(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("GetAssistantsWithLocaleTranslation", func(t *testing.T) { t.Run("GetAssistantsWithLocaleTranslation", func(t *testing.T) {
// Create assistant with i18n locales // Create assistant with i18n locales
@ -1950,13 +1902,12 @@ func TestGetAssistantsWithQueryFilter(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
// Create test assistants with different permission settings // Create test assistants with different permission settings
assistants := []types.AssistantModel{ assistants := []types.AssistantModel{
@ -2195,13 +2146,12 @@ func TestUpdateAssistant(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("UpdateSingleField", func(t *testing.T) { t.Run("UpdateSingleField", func(t *testing.T) {
// Create assistant // Create assistant
@ -3028,13 +2978,12 @@ func TestAssistantCompleteWorkflow(t *testing.T) {
test.Prepare(t, config.Conf) test.Prepare(t, config.Conf)
defer test.Clean() defer test.Clean()
store, err := NewXun(types.Setting{ store, err := xun.NewXun(types.Setting{
Connector: "default", Connector: "default",
}) })
if err != nil { if err != nil {
t.Fatalf("Failed to create store: %v", err) t.Fatalf("Failed to create store: %v", err)
} }
defer store.Close()
t.Run("CompleteWorkflow", func(t *testing.T) { t.Run("CompleteWorkflow", func(t *testing.T) {
// Step 1: Create multiple assistants // Step 1: Create multiple assistants

View file

@ -1,427 +1,39 @@
package xun package xun
import ( import (
"fmt"
"math"
"strings"
"time"
"github.com/yaoapp/yao/agent/i18n"
"github.com/yaoapp/yao/agent/store/types" "github.com/yaoapp/yao/agent/store/types"
) )
// UpdateChatTitle update the chat title // =============================================================================
func (conv *Xun) UpdateChatTitle(sid string, cid string, title string) error { // Chat Management
userID, err := conv.getUserID(sid) // =============================================================================
if err != nil {
return err
}
_, err = conv.newQueryChat(). // CreateChat creates a new chat session
Where("sid", userID). func (store *Xun) CreateChat(chat *types.Chat) error {
Where("chat_id", cid). // TODO: implement
Update(map[string]interface{}{ return nil
"title": title,
"updated_at": time.Now(),
})
return err
} }
// GetChat get the chat info and its history // GetChat retrieves a single chat by ID
func (conv *Xun) GetChat(sid string, cid string, locale ...string) (*types.ChatInfo, error) { func (store *Xun) GetChat(chatID string) (*types.Chat, error) {
// userID, err := conv.getUserID(sid) // TODO: implement
// if err != nil {
// return nil, err
// }
// Get chat info
qb := conv.newQueryChat().
Select("chat_id", "title", "assistant_id").
// Where("sid", userID).
Where("chat_id", cid)
row, err := qb.First()
if err != nil {
return nil, err
}
// Return nil if chat_id is nil (means no chat found)
if row.Get("chat_id") == nil {
return nil, nil return nil, nil
}
chat := map[string]interface{}{
"chat_id": row.Get("chat_id"),
"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
}
name := assistant.Get("name")
if len(locale) > 0 {
lang := strings.ToLower(locale[0])
name = i18n.Translate(assistantID.(string), lang, name).(string)
}
if assistant != nil {
chat["assistant_name"] = name
chat["assistant_avatar"] = assistant.Get("avatar")
}
}
// Get chat history with default filter (silent=false)
history, err := conv.GetHistory(sid, cid, locale...)
if err != nil {
return nil, err
}
return &types.ChatInfo{
Chat: chat,
History: history,
}, nil
} }
// GetChatWithFilter get the chat info and its history with filter options // UpdateChat updates chat fields
func (conv *Xun) GetChatWithFilter(sid string, cid string, filter types.ChatFilter, locale ...string) (*types.ChatInfo, error) { func (store *Xun) UpdateChat(chatID string, updates map[string]interface{}) error {
// userID, err := conv.getUserID(sid) // TODO: implement
// if err != nil { return nil
// return nil, err }
// }
// Get chat info // DeleteChat deletes a chat and its associated messages
qb := conv.newQueryChat(). func (store *Xun) DeleteChat(chatID string) error {
Select("chat_id", "title", "assistant_id"). // TODO: implement
// Where("sid", userID). return nil
Where("chat_id", cid) }
row, err := qb.First() // ListChats retrieves a paginated list of chats with optional grouping
if err != nil { func (store *Xun) ListChats(filter types.ChatFilter) (*types.ChatList, error) {
return nil, err // TODO: implement
}
// Return nil if chat_id is nil (means no chat found)
if row.Get("chat_id") == nil {
return nil, nil return nil, nil
}
chat := map[string]interface{}{
"chat_id": row.Get("chat_id"),
"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 with filter
history, err := conv.GetHistoryWithFilter(sid, cid, filter, locale...)
if err != nil {
return nil, err
}
return &types.ChatInfo{
Chat: chat,
History: history,
}, nil
}
// DeleteChat deletes a specific chat and its history
func (conv *Xun) DeleteChat(sid string, cid string) error {
userID, err := conv.getUserID(sid)
if err != nil {
return err
}
// Delete history records first
_, err = conv.newQuery().
Where("sid", userID).
Where("cid", cid).
Delete()
if err != nil {
return err
}
// Then delete the chat
_, err = conv.newQueryChat().
Where("sid", userID).
Where("chat_id", cid).
Limit(1).
Delete()
return err
}
// DeleteAllChats deletes all chats and their histories for a user
func (conv *Xun) DeleteAllChats(sid string) error {
userID, err := conv.getUserID(sid)
if err != nil {
return err
}
// Delete history records first
_, err = conv.newQuery().
Where("sid", userID).
Delete()
if err != nil {
return err
}
// Then delete all chats
_, err = conv.newQueryChat().
Where("sid", userID).
Delete()
return err
}
// GetChats get the chat list with grouping by date
func (conv *Xun) GetChats(sid string, filter types.ChatFilter, locale ...string) (*types.ChatGroupResponse, error) {
// Default behavior: exclude silent chats
if filter.Silent == nil {
silentFalse := false
filter.Silent = &silentFalse
}
return conv.getChatsWithFilter(sid, filter, locale...)
}
// getChatsWithFilter get the chats with filter options
func (conv *Xun) getChatsWithFilter(sid string, filter types.ChatFilter, locale ...string) (*types.ChatGroupResponse, error) {
// userID, err := conv.getUserID(sid)
// if err != nil {
// return nil, err
// }
// Set default values
if filter.Page <= 0 {
filter.Page = 1
}
if filter.PageSize <= 0 {
filter.PageSize = 20
}
if filter.Order == "" {
filter.Order = "desc"
}
// Get total count
qbCount := conv.newQueryChat()
// Where("sid", userID)
// Apply silent filter if provided
if filter.Silent != nil {
if *filter.Silent {
// Include all chats (both silent and non-silent)
} else {
// Only include non-silent chats
qbCount.Where("silent", false)
}
}
// Apply keyword filter if provided
if filter.Keywords != "" {
qbCount.Where("title", "like", fmt.Sprintf("%%%s%%", filter.Keywords))
}
total, err := qbCount.Count()
if err != nil {
return nil, err
}
// Calculate last page
lastPage := int(math.Ceil(float64(total) / float64(filter.PageSize)))
if lastPage < 1 {
lastPage = 1
}
// Get chats with pagination
qb := conv.newQueryChat().
Select("chat_id", "title", "assistant_id", "silent", "created_at", "updated_at")
// Where("sid", userID)
// Apply silent filter if provided
if filter.Silent != nil {
if *filter.Silent {
// Include all chats (both silent and non-silent)
} else {
// Only include non-silent chats
qb.Where("silent", false)
}
}
// Apply keyword filter if provided
if filter.Keywords != "" {
qb.Where("title", "like", fmt.Sprintf("%%%s%%", filter.Keywords))
}
// Apply pagination
offset := (filter.Page - 1) * filter.PageSize
qb.OrderBy("updated_at", filter.Order).
Offset(offset).
Limit(filter.PageSize)
rows, err := qb.Get()
if err != nil {
return nil, err
}
// 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)
lastWeekEnd := thisWeekStart.AddDate(0, 0, -1)
groups := map[string][]map[string]interface{}{
"Today": {},
"Yesterday": {},
"This Week": {},
"Last Week": {},
"Even Earlier": {},
}
// Collect assistant IDs to fetch their details
assistantIDs := []interface{}{}
for _, row := range rows {
if assistantID := row.Get("assistant_id"); assistantID != nil && assistantID != "" {
assistantIDs = append(assistantIDs, assistantID)
}
}
// Fetch assistant details
assistantMap := map[string]map[string]interface{}{}
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 {
name := assistant.Get("name")
if len(locale) > 0 {
lang := strings.ToLower(locale[0])
name = i18n.Translate(id.(string), lang, name).(string)
}
assistantMap[fmt.Sprintf("%v", id)] = map[string]interface{}{
"name": name,
"avatar": assistant.Get("avatar"),
}
}
}
}
for _, row := range rows {
chatID := row.Get("chat_id")
if chatID == nil || chatID == "" {
continue
}
chat := map[string]interface{}{
"chat_id": chatID,
"title": row.Get("title"),
"assistant_id": row.Get("assistant_id"),
"silent": row.Get("silent"),
}
// Add assistant details if available
if assistantID := row.Get("assistant_id"); assistantID != nil && assistantID != "" {
if assistant, ok := assistantMap[fmt.Sprintf("%v", assistantID)]; ok {
name := assistant["name"]
if len(locale) > 0 {
lang := strings.ToLower(locale[0])
name = i18n.Translate(assistantID.(string), lang, name).(string)
}
chat["assistant_name"] = name
chat["assistant_avatar"] = assistant["avatar"]
}
}
var dbDatetime = row.Get("updated_at")
if dbDatetime == nil {
dbDatetime = row.Get("created_at")
}
var createdAt time.Time
switch v := dbDatetime.(type) {
case time.Time:
createdAt = v
case string:
parsed, err := time.Parse("2006-01-02 15:04:05.999999-07:00", v)
if err != nil {
// Try alternative format
parsed, err = time.Parse(time.RFC3339, v)
if err != nil {
continue
}
}
createdAt = parsed
default:
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.Before(today):
groups["This Week"] = append(groups["This Week"], chat)
case createdDate.After(lastWeekStart) && createdDate.Before(lastWeekEnd.AddDate(0, 0, 1)):
groups["Last Week"] = append(groups["Last Week"], chat)
default:
groups["Even Earlier"] = append(groups["Even Earlier"], chat)
}
}
// Convert to ordered slice and apply i18n
result := []types.ChatGroup{}
for _, label := range []string{"Today", "Yesterday", "This Week", "Last Week", "Even Earlier"} {
if len(groups[label]) > 0 {
translatedLabel := label
if len(locale) > 0 {
lang := strings.ToLower(locale[0])
translatedLabel = i18n.TranslateGlobal(lang, label).(string)
}
result = append(result, types.ChatGroup{
Label: translatedLabel,
Chats: groups[label],
})
}
}
return &types.ChatGroupResponse{
Groups: result,
Page: filter.Page,
PageSize: filter.PageSize,
Total: total,
LastPage: lastPage,
}, nil
} }

View file

@ -1,322 +0,0 @@
package xun
import (
"fmt"
"strings"
"time"
"github.com/google/uuid"
jsoniter "github.com/json-iterator/go"
"github.com/yaoapp/yao/agent/i18n"
"github.com/yaoapp/yao/agent/store/types"
)
// GetHistory get the history
func (conv *Xun) GetHistory(sid string, cid string, locale ...string) ([]map[string]interface{}, error) {
// userID, err := conv.getUserID(sid)
// if err != nil {
// return nil, err
// }
qb := conv.newQuery().
Select("role", "name", "content", "context", "assistant_id", "assistant_name", "assistant_avatar", "mentions", "uid", "silent", "created_at", "updated_at").
// Where("sid", userID).
Where("cid", cid).
OrderBy("id", "desc")
// By default, exclude silent messages
qb.Where("silent", false)
if conv.setting.TTL > 0 {
qb.Where("expired_at", ">", time.Now())
}
limit := 20
if conv.setting.MaxSize > 0 {
limit = conv.setting.MaxSize
}
rows, err := qb.Limit(limit).Get()
if err != nil {
return nil, err
}
res := []map[string]interface{}{}
for _, row := range rows {
assistantName := row.Get("assistant_name")
assistantID := row.Get("assistant_id")
if len(locale) > 0 && assistantID != nil {
lang := strings.ToLower(locale[0])
assistantName = i18n.Translate(assistantID.(string), lang, assistantName).(string)
}
message := map[string]interface{}{
"role": row.Get("role"),
"name": row.Get("name"),
"content": row.Get("content"),
"context": row.Get("context"),
"assistant_id": row.Get("assistant_id"),
"assistant_name": assistantName,
"assistant_avatar": row.Get("assistant_avatar"),
"mentions": row.Get("mentions"),
"uid": row.Get("uid"),
"silent": row.Get("silent"),
"created_at": row.Get("created_at"),
"updated_at": row.Get("updated_at"),
}
res = append([]map[string]interface{}{message}, res...)
}
return res, nil
}
// SaveHistory save the history
func (conv *Xun) SaveHistory(sid string, messages []map[string]interface{}, cid string, context map[string]interface{}) error {
if cid == "" {
cid = uuid.New().String() // Generate a new UUID if cid is empty
}
userID, err := conv.getUserID(sid)
if err != nil {
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
}
}
// Get silent flag from context
var silent bool = false
var historyVisible bool = true
if context != nil {
if silentVal, ok := context["silent"]; ok {
switch v := silentVal.(type) {
case bool:
silent = v
case string:
silent = v == "true" || v == "1" || v == "yes"
case int:
silent = v != 0
case float64:
silent = v != 0
}
}
// Get history visible from context
if historyVisibleVal, ok := context["history_visible"]; ok {
switch v := historyVisibleVal.(type) {
case bool:
historyVisible = v
case string:
historyVisible = v == "true" || v == "1" || v == "yes"
case int:
historyVisible = v != 0
case float64:
historyVisible = v != 0
}
}
}
// First ensure chat record exists
exists, err := conv.newQueryChat().
Where("chat_id", cid).
Where("sid", userID).
Exists()
if err != nil {
return err
}
if !exists {
// Create new chat record
err = conv.newQueryChat().
Insert(map[string]interface{}{
"chat_id": cid,
"sid": userID,
"assistant_id": assistantID,
"silent": silent || historyVisible == false,
"created_at": time.Now(),
})
if err != nil {
return err
}
} else {
// Update assistant_id and silent if needed
_, err = conv.newQueryChat().
Where("chat_id", cid).
Where("sid", userID).
Update(map[string]interface{}{
"assistant_id": assistantID,
"silent": silent || historyVisible == false,
})
if err != nil {
return err
}
}
// Save message history
var expiredAt interface{} = nil
values := []map[string]interface{}{}
if conv.setting.TTL > 0 {
expiredAt = time.Now().Add(time.Duration(conv.setting.TTL) * time.Second)
}
now := time.Now()
for _, message := range messages {
// Type assertion safety checks
role, ok := message["role"].(string)
if !ok {
return fmt.Errorf("invalid role type in message: %v", message["role"])
}
content, ok := message["content"].(string)
if !ok {
return fmt.Errorf("invalid content type in message: %v", message["content"])
}
var contextRaw interface{} = nil
if context != nil {
contextRaw, err = jsoniter.MarshalToString(context)
if err != nil {
return err
}
}
// Process mentions if present
var mentionsRaw interface{} = nil
if mentions, ok := message["mentions"].([]interface{}); ok && len(mentions) > 0 {
mentionsRaw, err = jsoniter.MarshalToString(mentions)
if err != nil {
return err
}
}
value := map[string]interface{}{
"role": role,
"name": "",
"content": content,
"sid": userID,
"cid": cid,
"uid": userID,
"context": contextRaw,
"mentions": mentionsRaw,
"assistant_id": nil,
"assistant_name": nil,
"assistant_avatar": nil,
"silent": silent,
"created_at": now,
"updated_at": nil,
"expired_at": expiredAt,
}
if name, ok := message["name"].(string); ok {
value["name"] = name
}
// Add assistant fields if present
if assistantID, ok := message["assistant_id"].(string); ok {
value["assistant_id"] = assistantID
}
if assistantName, ok := message["assistant_name"].(string); ok {
value["assistant_name"] = assistantName
}
if assistantAvatar, ok := message["assistant_avatar"].(string); ok {
value["assistant_avatar"] = assistantAvatar
}
values = append(values, value)
}
err = conv.newQuery().Insert(values)
if err != nil {
return err
}
// Update Chat updated_at
_, err = conv.newQueryChat().
Where("chat_id", cid).
Where("sid", userID).
Update(map[string]interface{}{"updated_at": now})
if err != nil {
return err
}
return nil
}
// GetHistoryWithFilter get the history with filter options
func (conv *Xun) GetHistoryWithFilter(sid string, cid string, filter types.ChatFilter, locale ...string) ([]map[string]interface{}, error) {
userID, err := conv.getUserID(sid)
if err != nil {
return nil, err
}
qb := conv.newQuery().
Select("role", "name", "content", "context", "assistant_id", "assistant_name", "assistant_avatar", "mentions", "uid", "silent", "created_at", "updated_at").
Where("sid", userID).
Where("cid", cid).
OrderBy("id", "desc")
// Apply silent filter if provided, otherwise exclude silent messages by default
if filter.Silent != nil {
if *filter.Silent {
// Include all messages (both silent and non-silent)
} else {
// Only include non-silent messages
qb.Where("silent", false)
}
} else {
// Default behavior: exclude silent messages
qb.Where("silent", false)
}
if conv.setting.TTL > 0 {
qb.Where("expired_at", ">", time.Now())
}
limit := 20
if conv.setting.MaxSize > 0 {
limit = conv.setting.MaxSize
}
if filter.PageSize > 0 {
limit = filter.PageSize
}
// Apply pagination if provided
if filter.Page > 0 {
offset := (filter.Page - 1) * limit
qb.Offset(offset)
}
rows, err := qb.Limit(limit).Get()
if err != nil {
return nil, err
}
res := []map[string]interface{}{}
for _, row := range rows {
message := map[string]interface{}{
"role": row.Get("role"),
"name": row.Get("name"),
"content": row.Get("content"),
"context": row.Get("context"),
"assistant_id": row.Get("assistant_id"),
"assistant_name": row.Get("assistant_name"),
"assistant_avatar": row.Get("assistant_avatar"),
"mentions": row.Get("mentions"),
"uid": row.Get("uid"),
"silent": row.Get("silent"),
"created_at": row.Get("created_at"),
"updated_at": row.Get("updated_at"),
}
res = append([]map[string]interface{}{message}, res...)
}
return res, nil
}

View file

@ -0,0 +1,36 @@
package xun
import (
"github.com/yaoapp/yao/agent/store/types"
)
// =============================================================================
// Message Management
// =============================================================================
// SaveMessages batch saves messages for a chat
// This is the primary write method - messages are buffered during execution
// and batch-written at the end of a request
func (store *Xun) SaveMessages(chatID string, messages []*types.Message) error {
// TODO: implement
return nil
}
// GetMessages retrieves messages for a chat with filtering
func (store *Xun) GetMessages(chatID string, filter types.MessageFilter) ([]*types.Message, error) {
// TODO: implement
return nil, nil
}
// UpdateMessage updates a single message
func (store *Xun) UpdateMessage(messageID string, updates map[string]interface{}) error {
// TODO: implement
return nil
}
// DeleteMessages deletes specific messages from a chat
func (store *Xun) DeleteMessages(chatID string, messageIDs []string) error {
// TODO: implement
return nil
}

49
agent/store/xun/resume.go Normal file
View file

@ -0,0 +1,49 @@
package xun
import (
"github.com/yaoapp/yao/agent/store/types"
)
// =============================================================================
// Resume Management (only called on failure/interrupt)
// =============================================================================
// SaveResume batch saves resume records
// Only called when request is interrupted or failed
func (store *Xun) SaveResume(records []*types.Resume) error {
// TODO: implement
return nil
}
// GetResume retrieves all resume records for a chat
func (store *Xun) GetResume(chatID string) ([]*types.Resume, error) {
// TODO: implement
return nil, nil
}
// GetLastResume retrieves the last (most recent) resume record for a chat
func (store *Xun) GetLastResume(chatID string) (*types.Resume, error) {
// TODO: implement
return nil, nil
}
// GetResumeByStackID retrieves resume records for a specific stack
func (store *Xun) GetResumeByStackID(stackID string) ([]*types.Resume, error) {
// TODO: implement
return nil, nil
}
// GetStackPath returns the stack path from root to the given stack
// Returns: [root_stack_id, ..., current_stack_id]
func (store *Xun) GetStackPath(stackID string) ([]string, error) {
// TODO: implement
return nil, nil
}
// DeleteResume deletes all resume records for a chat
// Called after successful resume to clean up
func (store *Xun) DeleteResume(chatID string) error {
// TODO: implement
return nil
}

View file

@ -7,7 +7,6 @@ import (
jsoniter "github.com/json-iterator/go" jsoniter "github.com/json-iterator/go"
"github.com/yaoapp/gou/connector" "github.com/yaoapp/gou/connector"
"github.com/yaoapp/gou/model" "github.com/yaoapp/gou/model"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/xun/capsule" "github.com/yaoapp/xun/capsule"
"github.com/yaoapp/xun/dbal/query" "github.com/yaoapp/xun/dbal/query"
"github.com/yaoapp/xun/dbal/schema" "github.com/yaoapp/xun/dbal/schema"
@ -18,274 +17,155 @@ import (
// Xun implements the Store interface using a database backend. // Xun implements the Store interface using a database backend.
// It provides functionality for: // It provides functionality for:
// - Managing chat conversations and their message histories // - Managing chat sessions and their messages
// - Organizing chats with pagination and date-based grouping // - Organizing chats with pagination and date-based grouping
// - Handling chat metadata like titles and creation dates // - Handling chat metadata like titles and creation dates
// - Managing AI assistants with their configurations and metadata // - Managing AI assistants with their configurations and metadata
// - Supporting data expiration through TTL settings // - Managing resume records for recovery from interruptions
type Xun struct { type Xun struct {
query query.Query query query.Query
schema schema.Schema schema schema.Schema
setting types.Setting setting types.Setting
cleanTicker *time.Ticker
cleanStop chan bool
} }
// Public interface methods: // Public interface methods:
// //
// NewXun creates a new conversation instance with the given settings // NewXun creates a new store instance with the given settings
// GetChats retrieves a paginated list of chats grouped by date //
// GetChat retrieves a specific chat and its message history // Chat Management:
// GetChatWithFilter retrieves a specific chat with filter options // CreateChat creates a new chat session
// GetHistory retrieves the message history for a specific chat // GetChat retrieves a single chat by ID
// GetHistoryWithFilter retrieves the message history with filter options // UpdateChat updates chat fields
// SaveHistory saves new messages to a chat's historys // DeleteChat deletes a chat and its associated messages
// DeleteChat deletes a specific chat and its history // ListChats retrieves a paginated list of chats with optional grouping
// DeleteAllChats deletes all chats and their histories for a user //
// UpdateChatTitle updates the title of a specific chat // Message Management:
// SaveMessages batch saves messages for a chat
// GetMessages retrieves messages for a chat with filtering
// UpdateMessage updates a single message
// DeleteMessages deletes specific messages from a chat
//
// Resume Management:
// SaveResume batch saves resume records (only on failure/interrupt)
// GetResume retrieves all resume records for a chat
// GetLastResume retrieves the last resume record for a chat
// GetResumeByStackID retrieves resume records for a specific stack
// GetStackPath returns the stack path from root to the given stack
// DeleteResume deletes all resume records for a chat
//
// Assistant Management:
// SaveAssistant creates or updates an assistant // SaveAssistant creates or updates an assistant
// UpdateAssistant updates assistant fields
// DeleteAssistant deletes an assistant by assistant_id // DeleteAssistant deletes an assistant by assistant_id
// GetAssistants retrieves a paginated list of assistants with filtering // GetAssistants retrieves a paginated list of assistants with filtering
// GetAssistant retrieves a single assistant by assistant_id // GetAssistant retrieves a single assistant by assistant_id
// DeleteAssistants deletes assistants based on filter conditions // DeleteAssistants deletes assistants based on filter conditions
// GetAssistantTags retrieves all unique tags from assistants // GetAssistantTags retrieves all unique tags from assistants
// Close closes the store and releases any resources
// NewXun create a new xun store // NewXun create a new xun store
func NewXun(setting types.Setting) (types.Store, error) { func NewXun(setting types.Setting) (types.Store, error) {
conv := &Xun{setting: setting} store := &Xun{setting: setting}
if setting.Connector == "default" || setting.Connector == "" { if setting.Connector == "default" || setting.Connector == "" {
conv.query = capsule.Global.Query() store.query = capsule.Global.Query()
conv.schema = capsule.Global.Schema() store.schema = capsule.Global.Schema()
} else { } else {
conn, err := connector.Select(setting.Connector) conn, err := connector.Select(setting.Connector)
if err != nil { if err != nil {
return nil, fmt.Errorf("select store connector %s error: %s", setting.Connector, err.Error()) return nil, fmt.Errorf("select store connector %s error: %s", setting.Connector, err.Error())
} }
conv.query, err = conn.Query() store.query, err = conn.Query()
if err != nil { if err != nil {
return nil, fmt.Errorf("query store connector %s error: %s", setting.Connector, err.Error()) return nil, fmt.Errorf("query store connector %s error: %s", setting.Connector, err.Error())
} }
conv.schema, err = conn.Schema() store.schema, err = conn.Schema()
if err != nil { if err != nil {
return nil, err return nil, err
} }
} }
err := conv.initialize() return store, nil
if err != nil {
return nil, err
}
return conv, nil
} }
// Rename the following functions to start with lowercase letters to make them private: // =============================================================================
// Query Builders
// =============================================================================
func (conv *Xun) newQuery() query.Query { // newQueryChat creates a new query builder for the chat table
qb := conv.query.New() func (store *Xun) newQueryChat() query.Query {
qb.Table(conv.getHistoryTable()) qb := store.query.New()
qb.Table(store.getChatTable())
return qb return qb
} }
func (conv *Xun) newQueryChat() query.Query { // newQueryMessage creates a new query builder for the message table
qb := conv.query.New() func (store *Xun) newQueryMessage() query.Query {
qb.Table(conv.getChatTable()) qb := store.query.New()
qb.Table(store.getMessageTable())
return qb return qb
} }
func (conv *Xun) clean() { // newQueryResume creates a new query builder for the resume table
nums, err := conv.newQuery().Where("expired_at", "<=", time.Now()).Delete() func (store *Xun) newQueryResume() query.Query {
if err != nil { qb := store.query.New()
log.Error("Clean the conversation table error: %s", err.Error()) qb.Table(store.getResumeTable())
return return qb
}
if nums > 0 {
log.Trace("Clean the conversation table: %d", nums)
}
} }
// startAutoClean starts the automatic cleanup routine // newQueryAssistant creates a new query builder for the assistant table
func (conv *Xun) startAutoClean() { func (store *Xun) newQueryAssistant() query.Query {
if conv.cleanTicker != nil { qb := store.query.New()
conv.stopAutoClean() // Stop existing ticker if any qb.Table(store.getAssistantTable())
} return qb
conv.cleanTicker = time.NewTicker(1 * time.Hour) // Clean every hour
conv.cleanStop = make(chan bool)
go func() {
for {
select {
case <-conv.cleanTicker.C:
conv.clean()
case <-conv.cleanStop:
return
}
}
}()
log.Trace("Started automatic cleanup")
} }
// stopAutoClean stops the automatic cleanup routine // =============================================================================
func (conv *Xun) stopAutoClean() { // Table Name Getters
if conv.cleanTicker != nil { // =============================================================================
conv.cleanTicker.Stop()
conv.cleanTicker = nil
}
if conv.cleanStop != nil { // getChatTable returns the chat table name
close(conv.cleanStop) func (store *Xun) getChatTable() string {
conv.cleanStop = nil
}
log.Trace("Stopped automatic cleanup")
}
// Close stops the automatic cleanup and closes resources
func (conv *Xun) Close() error {
conv.stopAutoClean()
return nil
}
// Rename Init to initialize to avoid conflicts
func (conv *Xun) initialize() error {
// Start automatic cleanup if TTL is enabled
if conv.setting.TTL > 0 {
conv.startAutoClean()
}
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("cid", 200).Null().Index()
table.String("uid", 255).Null().Index()
table.String("role", 200).Null().Index()
table.String("name", 200).Null().Index()
table.Text("content").Null()
table.JSON("context").Null()
table.String("assistant_id", 200).Null().Index()
table.String("assistant_name", 200).Null()
table.String("assistant_avatar", 200).Null()
table.JSON("mentions").Null()
table.Boolean("silent").SetDefault(false).Index()
table.TimestampTz("created_at").SetDefaultRaw("CURRENT_TIMESTAMP").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", "cid", "uid", "role", "name", "content", "context", "assistant_id", "assistant_name", "assistant_avatar", "mentions", "silent", "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("assistant_id", 200).Null().Index()
table.String("sid", 255).Index()
table.Boolean("silent").SetDefault(false).Index()
table.TimestampTz("created_at").SetDefaultRaw("CURRENT_TIMESTAMP").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", "assistant_id", "sid", "silent", "created_at", "updated_at"}
for _, field := range fields {
if !tab.HasColumn(field) {
return fmt.Errorf("%s is required", field)
}
}
return nil
}
func (conv *Xun) getUserID(sid string) (string, error) {
// TODO: get the user id from the authentication system
return "guest", nil
}
func (conv *Xun) getHistoryTable() string {
m := model.Select("__yao.agent.history")
if m != nil && m.MetaData.Table.Name != "" {
return m.MetaData.Table.Name
}
return "__yao.agent.history"
}
func (conv *Xun) getChatTable() string {
m := model.Select("__yao.agent.chat") m := model.Select("__yao.agent.chat")
if m != nil && m.MetaData.Table.Name != "" { if m != nil && m.MetaData.Table.Name != "" {
return m.MetaData.Table.Name return m.MetaData.Table.Name
} }
return "__yao.agent.chat" return "agent_chat"
} }
func (conv *Xun) getAssistantTable() string { // getMessageTable returns the message table name
func (store *Xun) getMessageTable() string {
m := model.Select("__yao.agent.message")
if m != nil && m.MetaData.Table.Name != "" {
return m.MetaData.Table.Name
}
return "agent_message"
}
// getResumeTable returns the resume table name
func (store *Xun) getResumeTable() string {
m := model.Select("__yao.agent.resume")
if m != nil && m.MetaData.Table.Name != "" {
return m.MetaData.Table.Name
}
return "agent_resume"
}
// getAssistantTable returns the assistant table name
func (store *Xun) getAssistantTable() string {
m := model.Select("__yao.agent.assistant") m := model.Select("__yao.agent.assistant")
if m != nil && m.MetaData.Table.Name != "" { if m != nil && m.MetaData.Table.Name != "" {
return m.MetaData.Table.Name return m.MetaData.Table.Name
} }
return "__yao.agent.assistant" return "agent_assistant"
} }
// =============================================================================
// Utility Functions
// =============================================================================
// 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 (store *Xun) parseJSONFields(data map[string]interface{}, fields []string) {
for _, field := range fields { for _, field := range fields {
if val := data[field]; val != nil { if val := data[field]; val != nil {
if strVal, ok := val.(string); ok && strVal != "" { if strVal, ok := val.(string); ok && strVal != "" {
@ -299,7 +179,7 @@ func (conv *Xun) parseJSONFields(data map[string]interface{}, fields []string) {
} }
// GenerateAssistantID generates a random-looking 6-digit ID // GenerateAssistantID generates a random-looking 6-digit ID
func (conv *Xun) GenerateAssistantID() (string, error) { func (store *Xun) GenerateAssistantID() (string, error) {
maxAttempts := 10 // Maximum number of attempts to generate a unique ID maxAttempts := 10 // Maximum number of attempts to generate a unique ID
for i := 0; i < maxAttempts; i++ { for i := 0; i < maxAttempts; i++ {
// Generate a random number using timestamp and some bit operations // Generate a random number using timestamp and some bit operations
@ -308,8 +188,8 @@ func (conv *Xun) GenerateAssistantID() (string, error) {
hash := fmt.Sprintf("%06d", random) hash := fmt.Sprintf("%06d", random)
// Check if this ID already exists // Check if this ID already exists
exists, err := conv.query.New(). exists, err := store.query.New().
Table(conv.getAssistantTable()). Table(store.getAssistantTable()).
Where("assistant_id", hash). Where("assistant_id", hash).
Exists() Exists()