From 577c97aab0be03ee0ea3eb3f8b298c1123ec9d38 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=85=AC=E6=98=8E?= <83812544+Ed1s0nZ@users.noreply.github.com> Date: Wed, 22 Jul 2026 21:04:18 +0800 Subject: [PATCH] Add files via upload --- internal/database/conversation.go | 151 +++++++----------- internal/database/conversation_create_meta.go | 1 + internal/database/database.go | 16 ++ internal/database/project_stats.go | 8 +- 4 files changed, 79 insertions(+), 97 deletions(-) diff --git a/internal/database/conversation.go b/internal/database/conversation.go index 8241bdca..bf6d046d 100644 --- a/internal/database/conversation.go +++ b/internal/database/conversation.go @@ -22,6 +22,7 @@ type Conversation struct { ID string `json:"id"` Title string `json:"title"` ProjectID string `json:"projectId,omitempty"` + RoleName string `json:"roleName,omitempty"` Pinned bool `json:"pinned"` CreatedAt time.Time `json:"createdAt"` UpdatedAt time.Time `json:"updatedAt"` @@ -57,29 +58,30 @@ func (db *DB) CreateConversationWithWebshell(webshellConnectionID, title string, return nil, err } } + roleName := normalizeConversationRoleName(meta.RoleName) var err error wsID := strings.TrimSpace(webshellConnectionID) switch { case wsID != "" && projectID != "": _, err = db.Exec( - "INSERT INTO conversations (id, title, created_at, updated_at, webshell_connection_id, project_id) VALUES (?, ?, ?, ?, ?, ?)", - id, title, now, now, wsID, projectID, + "INSERT INTO conversations (id, title, created_at, updated_at, webshell_connection_id, project_id, role_name) VALUES (?, ?, ?, ?, ?, ?, ?)", + id, title, now, now, wsID, projectID, roleName, ) case wsID != "": _, err = db.Exec( - "INSERT INTO conversations (id, title, created_at, updated_at, webshell_connection_id) VALUES (?, ?, ?, ?, ?)", - id, title, now, now, wsID, + "INSERT INTO conversations (id, title, created_at, updated_at, webshell_connection_id, role_name) VALUES (?, ?, ?, ?, ?, ?)", + id, title, now, now, wsID, roleName, ) case projectID != "": _, err = db.Exec( - "INSERT INTO conversations (id, title, created_at, updated_at, project_id) VALUES (?, ?, ?, ?, ?)", - id, title, now, now, projectID, + "INSERT INTO conversations (id, title, created_at, updated_at, project_id, role_name) VALUES (?, ?, ?, ?, ?, ?)", + id, title, now, now, projectID, roleName, ) default: _, err = db.Exec( - "INSERT INTO conversations (id, title, created_at, updated_at) VALUES (?, ?, ?, ?)", - id, title, now, now, + "INSERT INTO conversations (id, title, created_at, updated_at, role_name) VALUES (?, ?, ?, ?, ?)", + id, title, now, now, roleName, ) } if err != nil { @@ -90,6 +92,7 @@ func (db *DB) CreateConversationWithWebshell(webshellConnectionID, title string, ID: id, Title: title, ProjectID: projectID, + RoleName: roleName, CreatedAt: now, UpdatedAt: now, } @@ -236,10 +239,11 @@ func (db *DB) GetConversation(id string) (*Conversation, error) { var pinned int var projectID sql.NullString + var roleName sql.NullString err := db.QueryRow( - "SELECT id, title, pinned, created_at, updated_at, project_id FROM conversations WHERE id = ?", + "SELECT id, title, pinned, created_at, updated_at, project_id, role_name FROM conversations WHERE id = ?", id, - ).Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID) + ).Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID, &roleName) if err != nil { if err == sql.ErrNoRows { return nil, fmt.Errorf("对话不存在") @@ -249,6 +253,9 @@ func (db *DB) GetConversation(id string) (*Conversation, error) { if projectID.Valid { conv.ProjectID = strings.TrimSpace(projectID.String) } + if roleName.Valid { + conv.RoleName = normalizeConversationRoleName(roleName.String) + } // 尝试多种时间格式解析 var err1, err2 error @@ -322,10 +329,11 @@ func (db *DB) GetConversationLite(id string) (*Conversation, error) { var pinned int var projectID sql.NullString + var roleName sql.NullString err := db.QueryRow( - "SELECT id, title, pinned, created_at, updated_at, project_id FROM conversations WHERE id = ?", + "SELECT id, title, pinned, created_at, updated_at, project_id, role_name FROM conversations WHERE id = ?", id, - ).Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID) + ).Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID, &roleName) if err != nil { if err == sql.ErrNoRows { return nil, fmt.Errorf("对话不存在") @@ -335,6 +343,9 @@ func (db *DB) GetConversationLite(id string) (*Conversation, error) { if projectID.Valid { conv.ProjectID = strings.TrimSpace(projectID.String) } + if roleName.Valid { + conv.RoleName = normalizeConversationRoleName(roleName.String) + } // 尝试多种时间格式解析 var err1, err2 error @@ -365,6 +376,26 @@ func (db *DB) GetConversationLite(id string) (*Conversation, error) { return &conv, nil } +func normalizeConversationRoleName(roleName string) string { + roleName = strings.TrimSpace(roleName) + if roleName == "" { + return "默认" + } + return roleName +} + +func (db *DB) SetConversationRoleName(id, roleName string) error { + roleName = normalizeConversationRoleName(roleName) + _, err := db.Exec( + "UPDATE conversations SET role_name = ?, updated_at = ? WHERE id = ?", + roleName, time.Now(), id, + ) + if err != nil { + return fmt.Errorf("更新对话角色失败: %w", err) + } + return nil +} + func conversationProjectIDColumn(alias string) string { if alias != "" { return alias + ".project_id" @@ -489,7 +520,7 @@ func (db *DB) ListConversations(limit, offset int, search, sortBy, projectID str where, args = appendConversationProjectFilter(where, args, projectID, "c") args = append(args, limit, offset) rows, err = db.Query( - `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id + `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id, c.role_name FROM conversations c`+where+` `+orderClause+` LIMIT ? OFFSET ?`, @@ -505,7 +536,7 @@ func (db *DB) ListConversations(limit, offset int, search, sortBy, projectID str } args = append(args, limit, offset) rows, err = db.Query( - "SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id FROM conversations"+where+" "+orderClause+" LIMIT ? OFFSET ?", + "SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id, role_name FROM conversations"+where+" "+orderClause+" LIMIT ? OFFSET ?", args..., ) } @@ -514,45 +545,7 @@ func (db *DB) ListConversations(limit, offset int, search, sortBy, projectID str return nil, fmt.Errorf("查询对话列表失败: %w", err) } defer rows.Close() - - var conversations []*Conversation - for rows.Next() { - var conv Conversation - var createdAt, updatedAt string - var pinned int - var projectID sql.NullString - - if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID); err != nil { - return nil, fmt.Errorf("扫描对话失败: %w", err) - } - if projectID.Valid { - conv.ProjectID = strings.TrimSpace(projectID.String) - } - - // 尝试多种时间格式解析 - var err1, err2 error - conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) - if err1 != nil { - conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt) - } - if err1 != nil { - conv.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) - } - - conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt) - if err2 != nil { - conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt) - } - if err2 != nil { - conv.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) - } - - conv.Pinned = pinned != 0 - - conversations = append(conversations, &conv) - } - - return conversations, nil + return scanConversationRows(rows) } func (db *DB) ListConversationsForAccess(limit, offset int, search, sortBy, projectID, userID, scope string) ([]*Conversation, error) { @@ -571,7 +564,7 @@ func (db *DB) ListConversationsForAccess(limit, offset int, search, sortBy, proj where, args = appendConversationAccessFilter(where, args, userID, scope, "c") args = append(args, limit, offset) rows, err = db.Query( - `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id + `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id, c.role_name FROM conversations c`+where+` `+orderClause+` LIMIT ? OFFSET ?`, args...) @@ -586,7 +579,7 @@ func (db *DB) ListConversationsForAccess(limit, offset int, search, sortBy, proj } args = append(args, limit, offset) rows, err = db.Query( - "SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id FROM conversations"+where+" "+orderClause+" LIMIT ? OFFSET ?", + "SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id, role_name FROM conversations"+where+" "+orderClause+" LIMIT ? OFFSET ?", args...) } if err != nil { @@ -603,12 +596,16 @@ func scanConversationRows(rows *sql.Rows) ([]*Conversation, error) { var createdAt, updatedAt string var pinned int var projectID sql.NullString - if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID); err != nil { + var roleName sql.NullString + if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID, &roleName); err != nil { return nil, fmt.Errorf("扫描对话失败: %w", err) } if projectID.Valid { conv.ProjectID = strings.TrimSpace(projectID.String) } + if roleName.Valid { + conv.RoleName = normalizeConversationRoleName(roleName.String) + } var err1, err2 error conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) if err1 != nil { @@ -668,7 +665,7 @@ func (db *DB) ListUngroupedConversations(limit, offset int, sortBy, projectID st where, args = appendConversationProjectFilter(where, args, projectID, "c") args = append(args, limit, offset) rows, err := db.Query( - `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id `+ + `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id, c.role_name `+ where+` `+orderClause+` LIMIT ? OFFSET ?`, @@ -678,43 +675,7 @@ func (db *DB) ListUngroupedConversations(limit, offset int, sortBy, projectID st return nil, fmt.Errorf("查询未分组对话失败: %w", err) } defer rows.Close() - - var conversations []*Conversation - for rows.Next() { - var conv Conversation - var createdAt, updatedAt string - var pinned int - var projectID sql.NullString - - if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID); err != nil { - return nil, fmt.Errorf("扫描对话失败: %w", err) - } - if projectID.Valid { - conv.ProjectID = strings.TrimSpace(projectID.String) - } - - var err1, err2 error - conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt) - if err1 != nil { - conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt) - } - if err1 != nil { - conv.CreatedAt, _ = time.Parse(time.RFC3339, createdAt) - } - - conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt) - if err2 != nil { - conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt) - } - if err2 != nil { - conv.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt) - } - - conv.Pinned = pinned != 0 - conversations = append(conversations, &conv) - } - - return conversations, rows.Err() + return scanConversationRows(rows) } func (db *DB) ListUngroupedConversationsForAccess(limit, offset int, sortBy, projectID, userID, scope string) ([]*Conversation, error) { @@ -728,7 +689,7 @@ func (db *DB) ListUngroupedConversationsForAccess(limit, offset int, sortBy, pro where, args = appendConversationAccessFilter(where, args, userID, scope, "c") args = append(args, limit, offset) rows, err := db.Query( - `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id `+ + `SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id, c.role_name `+ where+` `+orderClause+` LIMIT ? OFFSET ?`, diff --git a/internal/database/conversation_create_meta.go b/internal/database/conversation_create_meta.go index 8f94dc8e..08dc7ed0 100644 --- a/internal/database/conversation_create_meta.go +++ b/internal/database/conversation_create_meta.go @@ -5,6 +5,7 @@ type ConversationCreateMeta struct { Source string WebShellConnectionID string ProjectID string + RoleName string ClientIP string SessionHint string } diff --git a/internal/database/database.go b/internal/database/database.go index 4fa2b7ca..1f9ecc05 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -183,6 +183,7 @@ func (db *DB) initTables() error { title TEXT NOT NULL, created_at DATETIME NOT NULL, updated_at DATETIME NOT NULL, + role_name TEXT NOT NULL DEFAULT '默认', last_react_input TEXT, last_react_output TEXT );` @@ -1151,6 +1152,21 @@ func (db *DB) migrateConversationsTable() error { } } + // 检查 role_name 字段是否存在(对话绑定的业务角色,用于历史任务切换时恢复角色上下文) + err = db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('conversations') WHERE name='role_name'").Scan(&count) + if err != nil { + if _, addErr := db.Exec("ALTER TABLE conversations ADD COLUMN role_name TEXT NOT NULL DEFAULT '默认'"); addErr != nil { + errMsg := strings.ToLower(addErr.Error()) + if !strings.Contains(errMsg, "duplicate column") && !strings.Contains(errMsg, "already exists") { + db.logger.Warn("添加role_name字段失败", zap.Error(addErr)) + } + } + } else if count == 0 { + if _, err := db.Exec("ALTER TABLE conversations ADD COLUMN role_name TEXT NOT NULL DEFAULT '默认'"); err != nil { + db.logger.Warn("添加role_name字段失败", zap.Error(err)) + } + } + return nil } diff --git a/internal/database/project_stats.go b/internal/database/project_stats.go index b35e3787..2352309c 100644 --- a/internal/database/project_stats.go +++ b/internal/database/project_stats.go @@ -84,7 +84,7 @@ func (db *DB) ListConversationsByProjectID(projectID string, limit, offset int) limit = 100 } rows, err := db.Query( - `SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id + `SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id, role_name FROM conversations WHERE project_id = ? ORDER BY updated_at DESC LIMIT ? OFFSET ?`, projectID, limit, offset, ) @@ -99,12 +99,16 @@ func (db *DB) ListConversationsByProjectID(projectID string, limit, offset int) var createdAt, updatedAt string var pinned int var pid sql.NullString - if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &pid); err != nil { + var roleName sql.NullString + if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &pid, &roleName); err != nil { return nil, err } if pid.Valid { conv.ProjectID = strings.TrimSpace(pid.String) } + if roleName.Valid { + conv.RoleName = normalizeConversationRoleName(roleName.String) + } conv.CreatedAt = parseDBTime(createdAt) conv.UpdatedAt = parseDBTime(updatedAt) conv.Pinned = pinned != 0