DMP游戏内AI聊天的想法

8471 字
42 分钟
DMP游戏内AI聊天的想法
事情大概是这样的
闲得蛋疼发呆时,突然想到,关于饥荒这款游戏,里面是没有任何新手引导的,就等着你死一次长一智,死一次长一智
新手项快速成长,无非就是三种方式,多死、找大佬一起、查资料
如果是查资料的话,就要切出去,查看多个网页,我觉的挺割裂的,就想在游戏中做一个智能助理,专门解决游戏问题
说做就做,捋一下思路:
- 实时获取玩家聊天记录,需要毫秒级
- 判断聊天记录类型,玩家闲聊还是向AI提问
- 获取每个玩家的问题并分别保存,让AI有记忆且不混乱
- 调用大语言模型拿到答案
- 将答案发送到公屏
- 设置聊天记录定时清理,防止内存无限增长
有了这个思路,就开撸,下面是这一版游戏中的表现

符合预期的:
- 不会回答饥荒以外的东西
- 能够正确识别玩家的意图
- 聊天响应时间较快(小于1秒)
不符合预期的:
- 回答并不全对,回答的内容基本是靠AI训练时的数据,且污染严重
- 有很多数据没有,也是因为AI训练时,没有相关的游戏资料
那么接下来就是改进方案了:
- 构建一个权威的知识库(RAG),AI只能从知识库中获取资料并回答
- 知识库为MD文档
- 构建关键词索引和向量数据
- 搜索机制采用关键词搜索和向量搜索
- 查看是否配置了
embedding模型,有的话就向量搜索 - 没有
embedding就回退至关键词搜索
- 查看是否配置了
哈哈哈哈,说了这么多,卡在第1步了,我没有游戏的资料
我自己搞了点测试数据,试了一下
就比如
我:火炬如何制作
AI:火炬是游戏中可携带的光源,需要2个采下的草和2个树枝
网上爬是不可能的,没必要为了这个功能产生法务风险
TIP
期间也是发现了,某些人,用着开源的东西,却把自己有的东西,死死捏住不放,还阴阳怪气,跟狗一样
放弃了
下面是写好的代码,就留在这里,哪天有资料了,再捡起来
package aichat
import ( "context" "dst-management-platform-api/database/dao" "dst-management-platform-api/database/models" "sync")
// ================================================================// 导出类型// ================================================================
// Manager AI 对话服务管理器,是 aichat 包的核心入口。// 负责管理游戏内 AI 聊天的工作协程、Wiki 知识库搜索(关键词+向量)。type Manager struct { roomAISettingDao *dao.RoomAISettingDAO client *aiClient ctx context.Context cancel context.CancelFunc lifecycle sync.Mutex mu sync.Mutex workers map[int]*roomWorker active bool closed bool
keywordSearcher *keywordWikiSearcher embedSearcher *embeddingWikiSearcher embedSearcherMu sync.Mutex lastEmbedConfig string}
// EmbeddingConfig 向量嵌入模型配置type EmbeddingConfig struct { APIURL string APIKey string Model string Dimensions int}
// EmbeddingStats 向量索引统计信息type EmbeddingStats struct { TotalDocs int `json:"totalDocs"` TotalCategories int `json:"totalCategories"` VectorDim int `json:"vectorDim"` Categories map[string]int `json:"categories"`}
// NewManager 创建 AI 对话服务管理器func NewManager(roomAISettingDao *dao.RoomAISettingDAO, pluginDao *dao.PluginDAO) *Manager { return newManager(roomAISettingDao, pluginDao)}
// Start 启动所有已启用房间的 AI 对话监听func (m *Manager) Start() error { return m.start()}
// StopAll 停止所有房间的 AI 对话监听func (m *Manager) StopAll() { m.stopAll()}
// StopRoom 停止指定房间的 AI 对话监听func (m *Manager) StopRoom(roomID int) { m.lifecycle.Lock() defer m.lifecycle.Unlock() m.stopRoom(roomID)}
// Reload 重新加载指定房间的 AI 配置并重启监听func (m *Manager) Reload(roomID int) error { return m.reload(roomID)}
// Close 关闭 AI 对话服务,释放所有资源func (m *Manager) Close() { m.closeManager()}
// BuildEmbeddingIndex 手动构建向量索引。config 指定 Embedding API 配置,force 为 true 时强制全量重建。// 构建过程可能耗时较长。func (m *Manager) BuildEmbeddingIndex(config EmbeddingConfig, force bool) error { return m.buildEmbeddingIndex(config, force)}
// GetEmbeddingStats 获取向量索引统计信息func (m *Manager) GetEmbeddingStats(config EmbeddingConfig) EmbeddingStats { return m.getEmbeddingStats(config)}
// BuildKeywordIndex 手动构建关键词搜索索引。force 为 true 时强制重建。func (m *Manager) BuildKeywordIndex(force bool) error { return m.buildKeywordIndex(force)}
// UnloadKeywordIndex 释放关键词搜索索引占用的内存。下次 Search 时会自动从磁盘重新加载。func (m *Manager) UnloadKeywordIndex() { m.unloadKeywordIndex()}
// ValidateModelConfig 校验 AI 模型配置参数func ValidateModelConfig(config *models.AIModelConfig) error { return validateModelConfig(config)}
// ValidateRoomSetting 校验房间 AI 设置参数func ValidateRoomSetting(setting *models.RoomAISetting) error { return validateRoomSetting(setting)}package aichat
import ( "bytes" "context" "dst-management-platform-api/database/models" "dst-management-platform-api/utils" "encoding/json" "fmt" "io" "net/http" "net/url" "strings" "time")
type chatMessage struct { Role string `json:"role"` Content string `json:"content"`}
type chatCompletionRequest struct { Model string `json:"model"` Messages []chatMessage `json:"messages"` Temperature float64 `json:"temperature"` MaxTokens int `json:"max_tokens"` Stream bool `json:"stream"`}
type chatCompletionResponse struct { Choices []struct { Message chatMessage `json:"message"` } `json:"choices"` Error *struct { Message string `json:"message"` } `json:"error"`}
type aiClient struct { httpClient *http.Client}
func newClient() *aiClient { return &aiClient{ httpClient: &http.Client{ CheckRedirect: func(_ *http.Request, _ []*http.Request) error { return http.ErrUseLastResponse }, }, }}
func (c *aiClient) complete(ctx context.Context, config models.AIModelConfig, messages []chatMessage) (string, error) { requestBody := chatCompletionRequest{ Model: config.ChatModel, Messages: messages, Temperature: config.Temperature, MaxTokens: config.MaxTokens, Stream: false, } body, err := json.Marshal(requestBody) if err != nil { return "", fmt.Errorf("序列化大模型请求失败: %w", err) }
endpoint, err := chatCompletionsEndpoint(config.ChatBaseURL) if err != nil { return "", err } requestCtx, cancel := context.WithTimeout(ctx, time.Duration(config.RequestTimeoutSeconds)*time.Second) defer cancel()
req, err := http.NewRequestWithContext(requestCtx, http.MethodPost, endpoint, bytes.NewReader(body)) if err != nil { return "", fmt.Errorf("创建大模型请求失败: %w", err) } req.Header.Set("Content-Type", "application/json") req.Header.Set("User-Agent", fmt.Sprintf("DMP-AI/%s", utils.Version)) if config.ChatApiKey != "" { req.Header.Set("Authorization", "Bearer "+config.ChatApiKey) }
resp, err := c.httpClient.Do(req) if err != nil { return "", fmt.Errorf("请求大模型失败: %w", err) } defer resp.Body.Close()
responseBody, err := io.ReadAll(io.LimitReader(resp.Body, 2*1024*1024)) if err != nil { return "", fmt.Errorf("读取大模型响应失败: %w", err) }
var result chatCompletionResponse if len(responseBody) > 0 { if err = json.Unmarshal(responseBody, &result); err != nil { if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { return "", fmt.Errorf("大模型响应异常 HTTP %d", resp.StatusCode) } return "", fmt.Errorf("解析大模型响应失败: %w", err) } } if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { if result.Error != nil && result.Error.Message != "" { return "", fmt.Errorf("大模型响应异常 HTTP %d: %s", resp.StatusCode, result.Error.Message) } return "", fmt.Errorf("大模型响应异常 HTTP %d", resp.StatusCode) } if result.Error != nil && result.Error.Message != "" { return "", fmt.Errorf("大模型响应异常: %s", result.Error.Message) } if len(result.Choices) == 0 { return "", fmt.Errorf("大模型响应中没有候选答案") }
answer := strings.TrimSpace(result.Choices[0].Message.Content) if answer == "" { return "", fmt.Errorf("大模型返回了空答案") } return answer, nil}
func chatCompletionsEndpoint(baseURL string) (string, error) { u, err := url.Parse(strings.TrimSpace(baseURL)) if err != nil { return "", fmt.Errorf("解析大模型 Base URL 失败: %w", err) }
path := strings.TrimRight(u.Path, "/") if !strings.HasSuffix(path, "/chat/completions") { if path == "" { path = "/chat/completions" } else { path += "/chat/completions" } } u.Path = path return u.String(), nil}package aichat
import ( "dst-management-platform-api/database/models" "dst-management-platform-api/utils" "fmt" "strings" "unicode/utf8")
func validateModelConfig(config *models.AIModelConfig) error { config.ChatBaseURL = strings.TrimSpace(config.ChatBaseURL) config.ChatModel = strings.TrimSpace(config.ChatModel) config.EmbeddingModel = strings.TrimSpace(config.EmbeddingModel)
if config.ChatBaseURL == "" || !utils.IsValidURL(config.ChatBaseURL) { return fmt.Errorf("大模型 Base URL 不合法") } if config.ChatModel == "" { return fmt.Errorf("大模型名称不能为空") } if utf8.RuneCountInString(config.ChatModel) > 256 { return fmt.Errorf("大模型名称过长") } if len(config.ChatApiKey) > 16*1024 { return fmt.Errorf("API Key 过长") } if config.EmbeddingModel != "" { if utf8.RuneCountInString(config.EmbeddingModel) > 256 { return fmt.Errorf("嵌入模型名称过长") } if len(config.EmbeddingApiKey) > 16*1024 { return fmt.Errorf("嵌入 API Key 过长") } } if utf8.RuneCountInString(config.EmbeddingBaseURL) > 8000 { return fmt.Errorf("系统提示词过长") } if config.Temperature < 0 || config.Temperature > 2 { return fmt.Errorf("temperature 必须在 0 到 2 之间") } if config.MaxTokens <= 0 || config.MaxTokens > 32768 { return fmt.Errorf("maxTokens 必须在 1 到 32768 之间") } if config.RequestTimeoutSeconds <= 0 || config.RequestTimeoutSeconds > 300 { return fmt.Errorf("请求超时时间必须在 1 到 300 秒之间") }
return nil}
func validateRoomSetting(setting *models.RoomAISetting) error { if setting.RoomID <= 0 { return fmt.Errorf("房间 ID 不合法") } if strings.ContainsAny(setting.Prefix, "\r\n") || utf8.RuneCountInString(setting.Prefix) > 64 { return fmt.Errorf("AI 对话前缀不合法") } if setting.ContextMaxMessages < 2 || setting.ContextMaxMessages > 100 { return fmt.Errorf("上下文消息数量必须在 2 到 100 之间") } if setting.ContextTTLMinutes <= 0 || setting.ContextTTLMinutes > 10080 { return fmt.Errorf("上下文有效期必须在 1 到 10080 分钟之间") }
return ValidateModelConfig(&setting.AIModelConfig)}package aichat
import ( "context" "dst-management-platform-api/database/dao" "dst-management-platform-api/database/models" "dst-management-platform-api/dst" "dst-management-platform-api/logger" "errors" "fmt" "os" "strings" "time" "unicode/utf8"
"gorm.io/gorm")
const ( chatLogBufferSize = 256 maxQuestionRunes = 2000 maxReplyRunes = 180)
type chatSession struct { messages []chatMessage lastActive time.Time}
type roomWorker struct { cancel context.CancelFunc}
func newManager(roomAISettingDao *dao.RoomAISettingDAO, pluginDao *dao.PluginDAO) *Manager { ctx, cancel := context.WithCancel(context.Background())
aiManager := &Manager{ roomAISettingDao: roomAISettingDao, client: newClient(), ctx: ctx, cancel: cancel, workers: make(map[int]*roomWorker), keywordSearcher: newKeywordWikiSearcher(wikiPagesDir, wikiIndexFile), }
chatPlugin, err := pluginDao.GetPluginByPluginName(models.PluginChat) if err != nil { logger.Logger.Errorf("获取 %s 插件状态失败, err: %v", models.PluginChat, err) } else if chatPlugin.Status { if err = aiManager.start(); err != nil { logger.Logger.Errorf("启动游戏内 AI 对话服务失败, err: %v", err) } }
return aiManager}
func (m *Manager) start() error { m.lifecycle.Lock() if m.closed { m.lifecycle.Unlock() return context.Canceled } m.active = true m.lifecycle.Unlock()
settings, err := m.roomAISettingDao.ListEnabled() if err != nil { m.stopAll() return err } for _, setting := range settings { if err = m.reload(setting.RoomID); err != nil { logger.Logger.Errorf("启动房间 AI 对话监听失败, roomID: %d, err: %v", setting.RoomID, err) } } return nil}
func (m *Manager) reload(roomID int) error { m.lifecycle.Lock() defer m.lifecycle.Unlock() m.stopRoom(roomID) if !m.active || m.closed { return nil }
setting, err := m.roomAISettingDao.GetByRoomID(roomID) if errors.Is(err, gorm.ErrRecordNotFound) { return nil } if err != nil { return err } if !setting.Enabled { return nil } if err = validateRoomSetting(setting); err != nil { return err }
room, worlds, roomSetting, err := dao.FetchGameInfo(roomID) if err != nil { return err } game := dst.NewGameController(room, worlds, roomSetting, "zh") workerCtx, cancel := context.WithCancel(m.ctx)
m.mu.Lock() m.workers[roomID] = &roomWorker{cancel: cancel} m.mu.Unlock()
settingCopy := *setting go m.runRoom(workerCtx, game, settingCopy) return nil}
func (m *Manager) stopAll() { m.lifecycle.Lock() defer m.lifecycle.Unlock() m.active = false m.stopAllWorkers()}
func (m *Manager) stopRoom(roomID int) { m.mu.Lock() worker, ok := m.workers[roomID] if ok { delete(m.workers, roomID) } m.mu.Unlock()
if ok { worker.cancel() }}
func (m *Manager) closeManager() { m.lifecycle.Lock() defer m.lifecycle.Unlock() m.active = false m.closed = true m.cancel() m.stopAllWorkers()}
func (m *Manager) stopAllWorkers() { m.mu.Lock() workers := m.workers m.workers = make(map[int]*roomWorker) m.mu.Unlock() for _, worker := range workers { worker.cancel() }}
func (m *Manager) runRoom(ctx context.Context, game *dst.Game, setting models.RoomAISetting) { lines := make(chan string, chatLogBufferSize) go m.watchChatLog(ctx, game, setting.RoomID, lines)
sessions := make(map[string]*chatSession) cleanupTicker := time.NewTicker(time.Minute) defer cleanupTicker.Stop()
for { select { case <-ctx.Done(): return case <-cleanupTicker.C: cleanupSessions(sessions, time.Duration(setting.ContextTTLMinutes)*time.Minute) case line := <-lines: event, ok := parseChatEvent(line) if !ok { continue } question, ok := matchQuestion(setting.Prefix, event.Message) if !ok { continue } m.answer(ctx, game, setting, sessions, event, question) } }}
func (m *Manager) watchChatLog(ctx context.Context, game *dst.Game, roomID int, lines chan<- string) { for { err := game.TailChatLog(ctx, 0, lines) if ctx.Err() != nil { return } logger.Logger.Errorf("房间聊天日志监听异常, roomID: %d, err: %v", roomID, err)
timer := time.NewTimer(2 * time.Second) select { case <-ctx.Done(): timer.Stop() return case <-timer.C: } }}
// isEmbeddingConfigured 检查是否配置了嵌入模型func (m *Manager) isEmbeddingConfigured(setting models.RoomAISetting) bool { return setting.EmbeddingModel != "" && setting.ChatBaseURL != ""}
// getEmbeddingSearcher 获取或创建向量搜索引擎// 仅创建 searcher 实例并缓存,不自动构建索引(构建需通过 BuildEmbeddingIndex 手动触发)func (m *Manager) getEmbeddingSearcher(setting models.RoomAISetting) *embeddingWikiSearcher { // EmbeddingApiKey 如果单独配置了则优先使用,否则复用 ChatApiKey apiKey := setting.ChatApiKey if setting.EmbeddingApiKey != "" { apiKey = setting.EmbeddingApiKey } configKey := setting.ChatBaseURL + "|" + apiKey + "|" + setting.EmbeddingModel
m.embedSearcherMu.Lock() defer m.embedSearcherMu.Unlock()
// 配置未变,返回缓存的 searcher if m.embedSearcher != nil && m.lastEmbedConfig == configKey { return m.embedSearcher }
// 配置变更或首次调用 — 创建新的 searcher(不自动构建索引) if m.embedSearcher != nil { m.embedSearcher.stopIdleTimer() }
embedConfig := EmbeddingConfig{ APIURL: setting.ChatBaseURL, APIKey: apiKey, Model: setting.EmbeddingModel, Dimensions: 1024, } searcher := newEmbeddingWikiSearcher(wikiPagesDir, embedConfig)
m.embedSearcher = searcher m.lastEmbedConfig = configKey return m.embedSearcher}
// BuildEmbeddingIndex 手动构建向量索引。// config 指定 Embedding API 配置,force 为 true 时强制全量重建。// 构建过程可能耗时较长,建议在后台执行。func (m *Manager) buildEmbeddingIndex(config EmbeddingConfig, force bool) error { searcher := newEmbeddingWikiSearcher(wikiPagesDir, config)
if !force && !searcher.needsSetup() { logger.Logger.Infof("向量索引已存在,无需重建。使用 force=true 强制重建。") return nil }
logger.Logger.Infof("开始手动构建向量索引 (force=%v)...", force) if err := searcher.buildIndex(force); err != nil { return fmt.Errorf("构建向量索引失败: %w", err) }
// 更新缓存的 searcher m.embedSearcherMu.Lock() if m.embedSearcher != nil { m.embedSearcher.stopIdleTimer() } apiKey := config.APIKey configKey := config.APIURL + "|" + apiKey + "|" + config.Model m.embedSearcher = searcher m.lastEmbedConfig = configKey m.embedSearcherMu.Unlock()
logger.Logger.Infof("向量索引构建完成") return nil}
// GetEmbeddingStats 获取向量索引统计信息func (m *Manager) getEmbeddingStats(config EmbeddingConfig) EmbeddingStats { searcher := newEmbeddingWikiSearcher(wikiPagesDir, config) defer searcher.stopIdleTimer() return searcher.getStats()}
// BuildKeywordIndex 手动构建关键词搜索索引。// force 为 true 时强制重建,否则索引已存在则跳过。func (m *Manager) buildKeywordIndex(force bool) error { // 停止旧 searcher 的空闲计时器 if m.keywordSearcher != nil { m.keywordSearcher.stopIdleTimer() }
searcher := newKeywordWikiSearcher(wikiPagesDir, wikiIndexFile)
if !force { if _, err := os.Stat(wikiIndexFile); err == nil { if loadErr := searcher.load(); loadErr == nil { logger.Logger.Infof("关键词索引已存在,无需重建。使用 force=true 强制重建。") m.keywordSearcher = searcher return nil } } }
logger.Logger.Infof("开始构建关键词搜索索引 (force=%v)...", force) if err := searcher.buildIndex(force); err != nil { return fmt.Errorf("构建关键词索引失败: %w", err) }
m.keywordSearcher = searcher logger.Logger.Infof("关键词搜索索引构建完成") return nil}
// UnloadKeywordIndex 释放关键词搜索索引占用的内存。// 下次 Search 时会自动从磁盘重新加载。func (m *Manager) unloadKeywordIndex() { if m.keywordSearcher != nil { m.keywordSearcher.unload() }}
// searchWiki 搜索 Wiki 知识库,返回格式化的参考上下文func (m *Manager) searchWiki(setting models.RoomAISetting, question string) string { // 优先使用向量搜索 if m.isEmbeddingConfigured(setting) { if searcher := m.getEmbeddingSearcher(setting); searcher != nil { results, err := searcher.search(question, nil, 3, 0.3) if err != nil { logger.Logger.Warnf("向量搜索 Wiki 失败: %v,回退到关键词搜索", err) } else if len(results) > 0 { return formatWikiContext(results, maxContextTokens) } } }
// 回退到关键词搜索 if m.keywordSearcher != nil { results, err := m.keywordSearcher.search(question, 3) if err != nil { logger.Logger.Warnf("关键词搜索 Wiki 失败: %v", err) return "" } if len(results) > 0 { return formatWikiContext(results, maxContextTokens) } }
return ""}
// buildSystemPrompt 构建系统提示词(Wiki 上下文 + 用户设定的提示词)func buildSystemPrompt(setting models.RoomAISetting, wikiContext string) string { var parts []string
// 默认提示词 defaultPrompt := "你是饥荒联机版游戏内的 AI 助手。请根据以上参考文档,用中文回答玩家的问题。回答应简洁、准确,适合在游戏聊天框中显示,坚决不能使用使用 Markdown 格式,回答不能超过30个字。"
// Wiki 参考文档 if wikiContext != "" { parts = append(parts, wikiContext) parts = append(parts, "") parts = append(parts, defaultPrompt) return strings.Join(parts, "\n") }
// 用户设定的系统提示词(存储在 EmbeddingBaseURL 字段中) systemPrompt := strings.TrimSpace(setting.EmbeddingBaseURL) if systemPrompt != "" { return systemPrompt }
return defaultPrompt}
func (m *Manager) answer(ctx context.Context, game *dst.Game, setting models.RoomAISetting, sessions map[string]*chatSession, event chatEvent, question string) { now := time.Now() ttl := time.Duration(setting.ContextTTLMinutes) * time.Minute session := sessions[event.UID] if session == nil || now.Sub(session.lastActive) >= ttl { session = &chatSession{} sessions[event.UID] = session }
// 搜索 Wiki 知识库获取参考上下文 wikiContext := m.searchWiki(setting, question)
userMessage := chatMessage{Role: "user", Content: truncateRunes(question, maxQuestionRunes)} messages := make([]chatMessage, 0, len(session.messages)+3)
// 构建系统提示词(Wiki 上下文 + 用户设定的提示词) systemContent := buildSystemPrompt(setting, wikiContext) if systemContent != "" { messages = append(messages, chatMessage{Role: "system", Content: systemContent}) } messages = append(messages, session.messages...) messages = append(messages, userMessage)
answer, err := m.client.complete(ctx, setting.AIModelConfig, messages) if err != nil { if ctx.Err() != nil { return } logger.Logger.Errorf("游戏内 AI 回答失败, roomID: %d, uid: %s, err: %v", setting.RoomID, event.UID, err) _ = sendGameReply(game, event.Nickname, "暂时无法回答,请稍后再试。") return }
session.messages = append(session.messages, userMessage, chatMessage{Role: "assistant", Content: answer}) maxMessages := setting.ContextMaxMessages if maxMessages%2 != 0 { maxMessages-- } if len(session.messages) > maxMessages { session.messages = append([]chatMessage(nil), session.messages[len(session.messages)-maxMessages:]...) } session.lastActive = time.Now()
if err = sendGameReply(game, event.Nickname, answer); err != nil { logger.Logger.Errorf("发送游戏内 AI 回答失败, roomID: %d, uid: %s, err: %v", setting.RoomID, event.UID, err) }}
func cleanupSessions(sessions map[string]*chatSession, ttl time.Duration) { now := time.Now() for uid, session := range sessions { if now.Sub(session.lastActive) >= ttl { delete(sessions, uid) } }}
func matchQuestion(prefix, message string) (string, bool) { message = strings.TrimSpace(message) if prefix != "" { if !strings.HasPrefix(message, prefix) { return "", false } message = strings.TrimSpace(strings.TrimPrefix(message, prefix)) } return message, message != ""}
func sendGameReply(game *dst.Game, nickname, answer string) error { answer = strings.Join(strings.Fields(answer), " ") if answer == "" { return fmt.Errorf("AI 回答为空") }
prefix := fmt.Sprintf("[AI] %s: ", nickname) for _, part := range splitRunes(answer, maxReplyRunes-utf8.RuneCountInString(prefix)) { if err := game.SystemMsg(prefix + part); err != nil { return err } } return nil}
func splitRunes(value string, size int) []string { if size <= 0 { size = maxReplyRunes } runes := []rune(value) parts := make([]string, 0, (len(runes)+size-1)/size) for len(runes) > 0 { n := size if len(runes) < n { n = len(runes) } parts = append(parts, string(runes[:n])) runes = runes[n:] } return parts}
func truncateRunes(value string, size int) string { runes := []rune(value) if len(runes) <= size { return value } return string(runes[:size])}package aichat
import ( "regexp" "strings")
var sayLogPattern = regexp.MustCompile(`^\[[^\]]+\]:\s*\[Say\]\s*\(([^)]+)\)\s+(.+?):\s*(.*)$`)
type chatEvent struct { UID string Nickname string Message string}
func parseChatEvent(line string) (chatEvent, bool) { matches := sayLogPattern.FindStringSubmatch(strings.TrimSpace(line)) if len(matches) != 4 { return chatEvent{}, false }
event := chatEvent{ UID: strings.TrimSpace(matches[1]), Nickname: strings.TrimSpace(matches[2]), Message: strings.TrimSpace(matches[3]), } if event.UID == "" || event.Nickname == "" || event.Message == "" { return chatEvent{}, false } return event, true}package aichat
import ( "bytes" "dst-management-platform-api/logger" "dst-management-platform-api/utils" "encoding/json" "fmt" "io" "math" "net/http" "os" "path/filepath" "regexp" "sort" "strings" "sync" "time" "unicode/utf8")
// ========== 配置 ==========
const embeddingDir = utils.PluginAiChatSearchDataPath + "/embeddings"
// ========== 构建索引常量 ==========
const ( embedBatchSize = 10 embedMaxChars = 4000 embedRequestInterval = 300 * time.Millisecond embedMinBatchSize = 5 embedMinContentLength = 50 embeddingIdleTimeout = 5 * time.Minute)
// ========== 预处理正则 ==========
var ( embedMdTitleRe = regexp.MustCompile(`(?m)^#\s+.+$`) embedMdCategoryRe = regexp.MustCompile(`(?m)^\*\*分类\*\*:.*$`) embedMdHrRe = regexp.MustCompile(`(?m)^---$`) embedMdLinkRe = regexp.MustCompile(`\[([^\]]+)\]\([^)]+\)`) embedMdImgRe = regexp.MustCompile(`!\[.*?\]\(.*?\)`) embedMdFmtRe = regexp.MustCompile(`[*#>` + "`" + `|]`) embedWsRe = regexp.MustCompile(`\s+`) embedCatLinkRe = regexp.MustCompile(`\[([^\]]+)\]\(`))
// ========== 元数据结构 ==========
type wikiMeta struct { Title string `json:"title"` Categories []string `json:"categories"` Content string `json:"content"` ContentLength int `json:"content_length"` Filename string `json:"filename"`}
// ========== 向量搜索引擎 ==========
// EmbeddingWikiSearcher 基于向量相似度的 Wiki 搜索引擎type embeddingWikiSearcher struct { pagesDir string embeddingDir string apiURL string apiKey string model string dimensions int
httpClient *http.Client
mu sync.RWMutex metadata map[string]wikiMeta // filename -> meta embeddings map[string][]float64 // filename -> vector loaded bool idleTimer *time.Timer // 空闲自动释放计时器}
// newEmbeddingWikiSearcher 创建向量搜索引擎func newEmbeddingWikiSearcher(pagesDir string, config EmbeddingConfig) *embeddingWikiSearcher { if config.Dimensions <= 0 { config.Dimensions = 1024 }
return &embeddingWikiSearcher{ pagesDir: pagesDir, embeddingDir: embeddingDir, apiURL: strings.TrimRight(config.APIURL, "/"), apiKey: config.APIKey, model: config.Model, dimensions: config.Dimensions, httpClient: &http.Client{ Timeout: 120 * time.Second, }, }}
// NeedsSetup 是否需要构建索引func (s *embeddingWikiSearcher) needsSetup() bool { s.mu.RLock() defer s.mu.RUnlock()
if !s.loaded { s.mu.RUnlock() s.load() s.mu.RLock() } return len(s.embeddings) == 0}
// Search 向量语义搜索func (s *embeddingWikiSearcher) search(query string, categories []string, maxResults int, minScore float64) ([]wikiResult, error) { s.mu.RLock() if !s.loaded { s.mu.RUnlock() s.load() s.mu.RLock() } embeddings := s.embeddings metadata := s.metadata s.mu.RUnlock()
// 每次搜索重置空闲计时器 s.mu.Lock() s.resetIdleTimer() s.mu.Unlock()
if len(embeddings) == 0 { return nil, fmt.Errorf("向量索引为空") }
// 对查询文本做 embedding queryVector, err := s.embedSingle(query) if err != nil { return nil, fmt.Errorf("查询文本 embedding 失败: %w", err) }
// 计算余弦相似度 type scoredDoc struct { filename string score float64 } var scores []scoredDoc
for filename, docVec := range embeddings { meta, ok := metadata[filename] if !ok || meta.ContentLength < embedMinContentLength { continue } if len(categories) > 0 { catSet := make(map[string]struct{}, len(categories)) for _, c := range categories { catSet[c] = struct{}{} } hasCat := false for _, c := range meta.Categories { if _, ok := catSet[c]; ok { hasCat = true break } } if !hasCat { continue } }
sim := cosineSimilarity(queryVector, docVec) if sim >= minScore { scores = append(scores, scoredDoc{filename, sim}) } }
// 排序取 Top N sort.Slice(scores, func(i, j int) bool { return scores[i].score > scores[j].score })
results := make([]wikiResult, 0, maxResults) for i := 0; i < len(scores) && i < maxResults; i++ { meta := metadata[scores[i].filename] results = append(results, wikiResult{ Title: meta.Title, Filename: meta.Filename, Categories: meta.Categories, Content: meta.Content, ContentLen: meta.ContentLength, Score: roundFloat(scores[i].score, 4), }) }
return results, nil}
// ========== 内部方法 ==========
func (s *embeddingWikiSearcher) load() { s.mu.Lock() defer s.mu.Unlock() s.loadLocked()}
// loadLocked 不加锁地从磁盘加载(调用方需持有 s.mu 写锁)func (s *embeddingWikiSearcher) loadLocked() { if s.loaded { return }
metadataPath := filepath.Join(s.embeddingDir, "metadata.json") embeddingsPath := filepath.Join(s.embeddingDir, "embeddings.json")
if data, err := os.ReadFile(metadataPath); err == nil { err = json.Unmarshal(data, &s.metadata) if err != nil { logger.Logger.Errorf("载入向量模型失败: %v", err) } } if s.metadata == nil { s.metadata = make(map[string]wikiMeta) }
if data, err := os.ReadFile(embeddingsPath); err == nil { err = json.Unmarshal(data, &s.embeddings) if err != nil { logger.Logger.Errorf("载入向量模型失败: %v", err) } } if s.embeddings == nil { s.embeddings = make(map[string][]float64) }
s.loaded = true s.resetIdleTimer() logger.Logger.Infof("已加载向量索引: %d 篇文档", len(s.embeddings))}
// unload 释放向量索引占用的内存。下次 search 时会自动从磁盘重新加载。func (s *embeddingWikiSearcher) unload() { if s.embeddings != nil { logger.Logger.Infof("释放向量索引内存 (%d 篇文档)", len(s.embeddings)) } s.metadata = nil s.embeddings = nil s.loaded = false s.stopIdleTimer()}
// resetIdleTimer 重置空闲计时器。调用方需持有 s.mu 写锁。func (s *embeddingWikiSearcher) resetIdleTimer() { if s.idleTimer != nil { s.idleTimer.Stop() } s.idleTimer = time.AfterFunc(embeddingIdleTimeout, func() { s.mu.Lock() //logger.Logger.Infof("向量索引空闲超过 %v,自动释放内存", embeddingIdleTimeout) s.unload() s.mu.Unlock() })}
// stopIdleTimer 停止空闲计时器func (s *embeddingWikiSearcher) stopIdleTimer() { if s.idleTimer != nil { s.idleTimer.Stop() s.idleTimer = nil }}
func (s *embeddingWikiSearcher) embedSingle(text string) ([]float64, error) { vectors, err := s.embedBatch([]string{text}) if err != nil { return nil, err } if len(vectors) == 0 { return nil, fmt.Errorf("embedding 返回为空") } return vectors[0], nil}
func (s *embeddingWikiSearcher) embedBatch(texts []string) ([][]float64, error) { // 构建请求 reqBody := map[string]interface{}{ "model": s.model, "input": texts, "encoding_format": "float", } if s.dimensions > 0 { reqBody["dimensions"] = s.dimensions }
body, err := json.Marshal(reqBody) if err != nil { return nil, fmt.Errorf("序列化 embedding 请求失败: %w", err) }
endpoint := s.apiURL + "/embeddings" req, err := http.NewRequest(http.MethodPost, endpoint, bytes.NewReader(body)) if err != nil { return nil, fmt.Errorf("创建 embedding 请求失败: %w", err) } req.Header.Set("Content-Type", "application/json") if s.apiKey != "" { req.Header.Set("Authorization", "Bearer "+s.apiKey) }
resp, err := s.httpClient.Do(req) if err != nil { return nil, fmt.Errorf("embedding API 请求失败: %w", err) } defer resp.Body.Close()
respBody, err := io.ReadAll(io.LimitReader(resp.Body, 32*1024*1024)) if err != nil { return nil, fmt.Errorf("读取 embedding 响应失败: %w", err) }
if resp.StatusCode != http.StatusOK { return nil, fmt.Errorf("embedding API 返回 HTTP %d: %s", resp.StatusCode, string(respBody)) }
var result struct { Data []struct { Embedding []float64 `json:"embedding"` Index int `json:"index"` } `json:"data"` Error *struct { Message string `json:"message"` } `json:"error"` } if err := json.Unmarshal(respBody, &result); err != nil { return nil, fmt.Errorf("解析 embedding 响应失败: %w", err) } if result.Error != nil && result.Error.Message != "" { return nil, fmt.Errorf("embedding API 错误: %s", result.Error.Message) }
// 按 index 排序 sort.Slice(result.Data, func(i, j int) bool { return result.Data[i].Index < result.Data[j].Index })
vectors := make([][]float64, len(result.Data)) for i, d := range result.Data { vectors[i] = d.Embedding } return vectors, nil}
func cosineSimilarity(a, b []float64) float64 { if len(a) != len(b) || len(a) == 0 { return 0 }
var dot, normA, normB float64 for i := range a { dot += a[i] * b[i] normA += a[i] * a[i] normB += b[i] * b[i] }
if normA == 0 || normB == 0 { return 0 }
return dot / (math.Sqrt(normA) * math.Sqrt(normB))}
// ================================================================// 构建向量索引// ================================================================
// save 不加锁保存(调用方需持有 s.mu 读锁或写锁)func (s *embeddingWikiSearcher) save() error { if err := os.MkdirAll(s.embeddingDir, 0755); err != nil { return fmt.Errorf("创建 embedding 目录失败: %w", err) }
metadataPath := filepath.Join(s.embeddingDir, "metadata.json") embeddingsPath := filepath.Join(s.embeddingDir, "embeddings.json")
data, err := json.Marshal(s.metadata) if err != nil { return fmt.Errorf("序列化 metadata 失败: %w", err) } if err := os.WriteFile(metadataPath, data, 0644); err != nil { return fmt.Errorf("保存 metadata 失败: %w", err) }
data, err = json.Marshal(s.embeddings) if err != nil { return fmt.Errorf("序列化 embeddings 失败: %w", err) } if err := os.WriteFile(embeddingsPath, data, 0644); err != nil { return fmt.Errorf("保存 embeddings 失败: %w", err) }
return nil}
// BuildIndex 构建向量索引。只跑一次(或 force=true 强制重建)。//// 流程:// - 扫描所有 .md 页面// - 文本预处理// - 分批调用 Embedding API// - 每批完成后保存(支持断点续传)// - 连接错误自动重试,严重时自动降级 batch_sizefunc (s *embeddingWikiSearcher) buildIndex(force bool) error { s.mu.Lock() defer s.mu.Unlock()
// 确保已加载已有数据 s.loadLocked()
// 扫描所有 .md 页面 entries, err := os.ReadDir(s.pagesDir) if err != nil { return fmt.Errorf("扫描 Wiki 页面目录失败: %w", err) }
var mdFiles []string for _, entry := range entries { if !entry.IsDir() && filepath.Ext(entry.Name()) == ".md" { mdFiles = append(mdFiles, filepath.Join(s.pagesDir, entry.Name())) } } sort.Strings(mdFiles) total := len(mdFiles)
// 增量构建 vs 全量重建 if !force { var newFiles []string for _, f := range mdFiles { if _, ok := s.embeddings[filepath.Base(f)]; !ok { newFiles = append(newFiles, f) } } if len(newFiles) == 0 { logger.Logger.Infof("所有 %d 篇文档已有向量,无需重建。用 force=true 强制重建。", total) return nil } logger.Logger.Infof("增量构建: %d/%d 篇新文档需要 embedding", len(newFiles), total) mdFiles = newFiles } else { s.embeddings = make(map[string][]float64) s.metadata = make(map[string]wikiMeta) }
// 分批处理 currentBatchSize := embedBatchSize batches, totalBatches := makeBatches(mdFiles, currentBatchSize)
logger.Logger.Infof("共 %d 篇文档,分 %d 批 (batch_size=%d)", len(mdFiles), totalBatches, currentBatchSize) logger.Logger.Infof("API: %s 模型: %s", s.apiURL, s.model)
var failedFiles []string consecutiveFailures := 0 batchIdx := 0
for batchIdx < totalBatches { batchFiles := batches[batchIdx]
// 准备这批的文本 texts, metas, prepErr := s.prepareBatch(batchFiles) if prepErr != nil { logger.Logger.Errorf("准备第 %d 批文档失败: %v", batchIdx+1, prepErr) for _, f := range batchFiles { failedFiles = append(failedFiles, filepath.Base(f)) } batchIdx++ continue }
// 调用 API vectors, apiErr := s.embedBatch(texts)
if apiErr == nil { // 成功 for i, fp := range batchFiles { s.embeddings[filepath.Base(fp)] = vectors[i] } for k, v := range metas { s.metadata[k] = v } consecutiveFailures = 0
if saveErr := s.save(); saveErr != nil { logger.Logger.Errorf("保存 embedding 失败: %v", saveErr) } logger.Logger.Infof("第 %d/%d 批完成 (%d 篇)", batchIdx+1, totalBatches, len(s.embeddings))
if batchIdx < totalBatches-1 { time.Sleep(embedRequestInterval) } batchIdx++ continue }
// --- API 调用失败,错误恢复 --- consecutiveFailures++ errorStr := strings.ToLower(apiErr.Error()) logger.Logger.Errorf("第 %d/%d 批失败: %v", batchIdx+1, totalBatches, apiErr)
// batch size 超限 → 立刻调整为 10 if strings.Contains(errorStr, "batch size") && currentBatchSize > 10 { currentBatchSize = 10 logger.Logger.Infof("[自动调整] batch_size 调整为 %d(API 限制)", currentBatchSize) batches, totalBatches = rebuildBatches(mdFiles, s.embeddings, currentBatchSize) batchIdx = 0 consecutiveFailures = 0 continue }
// 连续失败 3 次,降级 batch_size if consecutiveFailures >= 3 && currentBatchSize > embedMinBatchSize { currentBatchSize = max(embedMinBatchSize, currentBatchSize/2) logger.Logger.Infof("[自动调整] batch_size 降为 %d,重新分批...", currentBatchSize) batches, totalBatches = rebuildBatches(mdFiles, s.embeddings, currentBatchSize) batchIdx = 0 consecutiveFailures = 0 continue }
// 把这一批拆成单篇重试 if currentBatchSize > 1 && len(batchFiles) > 1 { logger.Logger.Infof("尝试逐篇处理这一批...") for _, fp := range batchFiles { singleTexts, singleMetas, prepErr := s.prepareBatch([]string{fp}) if prepErr != nil { logger.Logger.Errorf("[跳过] %s: %v", filepath.Base(fp), prepErr) failedFiles = append(failedFiles, filepath.Base(fp)) continue } vec, embErr := s.embedBatch(singleTexts) if embErr != nil { logger.Logger.Errorf("[跳过] %s: %v", filepath.Base(fp), embErr) failedFiles = append(failedFiles, filepath.Base(fp)) continue } s.embeddings[filepath.Base(fp)] = vec[0] for k, v := range singleMetas { s.metadata[k] = v } time.Sleep(500 * time.Millisecond) } if saveErr := s.save(); saveErr != nil { logger.Logger.Errorf("保存 embedding 失败: %v", saveErr) } batchIdx++ consecutiveFailures = 0 continue }
// 单篇也失败,跳过 logger.Logger.Errorf("[跳过] 这 %d 篇 embed 失败", len(batchFiles)) for _, f := range batchFiles { failedFiles = append(failedFiles, filepath.Base(f)) } if saveErr := s.save(); saveErr != nil { logger.Logger.Errorf("保存 embedding 失败: %v", saveErr) } batchIdx++ }
// 完成 s.resetIdleTimer() logger.Logger.Infof("索引构建完成! 共 %d 篇文档。", len(s.embeddings)) if len(failedFiles) > 0 { logger.Logger.Warnf("跳过 %d 篇:", len(failedFiles)) for i, f := range failedFiles { if i >= 10 { logger.Logger.Warnf("... 等 %d 篇", len(failedFiles)-10) break } logger.Logger.Warnf(" - %s", f) } }
return nil}
// ========== 批次准备 ==========
// prepareBatch 批量读取并预处理文档func (s *embeddingWikiSearcher) prepareBatch(files []string) (texts []string, metas map[string]wikiMeta, err error) { texts = make([]string, 0, len(files)) metas = make(map[string]wikiMeta, len(files))
for _, fp := range files { content, readErr := os.ReadFile(fp) if readErr != nil { return nil, nil, fmt.Errorf("读取 %s 失败: %w", filepath.Base(fp), readErr) } rawContent := string(content)
title := extractWikiTitle(rawContent, filepath.Base(fp)) categories := extractWikiCategories(rawContent) cleanText := preprocessWikiDoc(rawContent, title) strippedContent := stripWikiMeta(rawContent)
texts = append(texts, cleanText) metas[filepath.Base(fp)] = wikiMeta{ Title: title, Categories: categories, Content: strippedContent, ContentLength: utf8.RuneCountInString(cleanText), Filename: filepath.Base(fp), } }
return texts, metas, nil}
// ========== 预处理 ==========
// extractWikiTitle 提取文档标题(markdown 一级标题)func extractWikiTitle(rawContent, fallback string) string { match := embedMdTitleRe.FindStringSubmatch(rawContent) if len(match) >= 1 { return strings.TrimSpace(strings.TrimPrefix(match[0], "# ")) } // 回退:文件名去扩展名 name := fallback if ext := filepath.Ext(name); ext != "" { name = strings.TrimSuffix(name, ext) } return name}
// extractWikiCategories 提取文档分类func extractWikiCategories(rawContent string) []string { match := embedMdCategoryRe.FindStringSubmatch(rawContent) if len(match) < 1 { return nil } linkMatches := embedCatLinkRe.FindAllStringSubmatch(match[0], -1) cats := make([]string, 0, len(linkMatches)) for _, m := range linkMatches { if len(m) >= 2 { cats = append(cats, m[1]) } } return cats}
// preprocessWikiDoc 预处理文档文本,使其适合 embedding//// 处理步骤:// - 去 markdown 语法// - 标题前置(标题是最重要的语义信号)// - 截断到 embedMaxCharsfunc preprocessWikiDoc(rawContent, title string) string { text := rawContent
// 去标题行、分类行、水平线 text = embedMdTitleRe.ReplaceAllString(text, "") text = embedMdCategoryRe.ReplaceAllString(text, "") text = embedMdHrRe.ReplaceAllString(text, "")
// 去 markdown 格式字符 text = embedMdLinkRe.ReplaceAllString(text, "$1") // [text](url) → text text = embedMdImgRe.ReplaceAllString(text, " ") text = embedMdFmtRe.ReplaceAllString(text, " ") text = embedWsRe.ReplaceAllString(text, " ") text = strings.TrimSpace(text)
// 标题前置 text = fmt.Sprintf("标题: %s\n正文: %s", title, text)
// 截断 runes := []rune(text) if len(runes) > embedMaxChars { text = string(runes[:embedMaxChars]) }
return text}
// stripWikiMeta 去除元数据行,保留纯正文func stripWikiMeta(rawContent string) string { text := rawContent text = embedMdTitleRe.ReplaceAllString(text, "") text = embedMdCategoryRe.ReplaceAllString(text, "") text = embedMdHrRe.ReplaceAllString(text, "") return strings.TrimSpace(text)}
// ========== 辅助函数 ==========
func makeBatches(files []string, size int) (batches [][]string, total int) { for i := 0; i < len(files); i += size { end := i + size if end > len(files) { end = len(files) } batches = append(batches, files[i:end]) } return batches, len(batches)}
func rebuildBatches(allFiles []string, embedded map[string][]float64, size int) ([][]string, int) { var remaining []string for _, f := range allFiles { if _, ok := embedded[filepath.Base(f)]; !ok { remaining = append(remaining, f) } } logger.Logger.Infof("剩余 %d 篇,分为 %d 批", len(remaining), (len(remaining)+size-1)/size) return makeBatches(remaining, size)}
// ================================================================// 索引统计// ================================================================
// GetStats 获取索引统计信息func (s *embeddingWikiSearcher) getStats() EmbeddingStats { s.mu.RLock() defer s.mu.RUnlock()
if !s.loaded { s.mu.RUnlock() s.load() s.mu.RLock() }
stats := EmbeddingStats{ TotalDocs: len(s.embeddings), Categories: make(map[string]int), }
for _, meta := range s.metadata { for _, c := range meta.Categories { stats.Categories[c]++ } } stats.TotalCategories = len(stats.Categories)
if len(s.embeddings) > 0 { for _, v := range s.embeddings { stats.VectorDim = len(v) break } }
return stats}package aichat
import ( "dst-management-platform-api/logger" "dst-management-platform-api/utils" "encoding/json" "fmt" "os" "path/filepath" "regexp" "sort" "strings" "sync" "time" "unicode" "unicode/utf8")
// ========== 配置 ==========
const ( wikiPagesDir = utils.PluginAiChatWikiPath + "/pages" wikiIndexFile = utils.PluginAiChatSearchDataPath + "/search_index.json" keywordIndexIdleTimeout = 5 * time.Minute)
// ========== 数据结构 ==========
// wikiResult 单条搜索结果type wikiResult struct { Title string `json:"title"` Filename string `json:"filename"` Categories []string `json:"categories"` Content string `json:"content"` ContentLen int `json:"content_length"` Score float64 `json:"score"`}
type wikiPageInfo struct { Title string `json:"title"` Categories []string `json:"categories"` Content string `json:"content"` ContentLength int `json:"content_length"` Links []string `json:"links"` Filename string `json:"filename"` Terms []string `json:"terms"`}
type wikiIndex struct { Pages map[string]wikiPageInfo `json:"pages"` CategoryIndex map[string][]string `json:"category_index"` TermIndex map[string]map[string]float64 `json:"term_index"`}
// ========== 关键词搜索引擎 ==========
// KeywordWikiSearcher 基于关键词的 Wiki 搜索引擎type keywordWikiSearcher struct { pagesDir string indexPath string
mu sync.RWMutex index *wikiIndex idleTimer *time.Timer // 空闲自动释放计时器}
// NewKeywordWikiSearcher 创建关键词搜索引擎func newKeywordWikiSearcher(pagesDir, indexPath string) *keywordWikiSearcher { return &keywordWikiSearcher{ pagesDir: pagesDir, indexPath: indexPath, }}
// Load 从磁盘加载搜索索引func (s *keywordWikiSearcher) load() error { s.mu.Lock() defer s.mu.Unlock() return s.loadLocked()}
// Unload 释放搜索索引占用的内存。下次 Search 时会自动重新从磁盘加载。func (s *keywordWikiSearcher) unload() { s.mu.Lock() defer s.mu.Unlock() if s.index != nil { logger.Logger.Infof("释放关键词搜索索引内存 (%d 篇文档)", len(s.index.Pages)) s.index = nil } s.stopIdleTimer()}
// IsLoaded 检查索引是否已加载到内存func (s *keywordWikiSearcher) isLoaded() bool { s.mu.RLock() defer s.mu.RUnlock() return s.index != nil}
// resetIdleTimer 重置空闲计时器。调用方需持有 s.mu 写锁。// 每次 Search 调用时重置,超时后自动 Unload。func (s *keywordWikiSearcher) resetIdleTimer() { if s.idleTimer != nil { s.idleTimer.Stop() } s.idleTimer = time.AfterFunc(keywordIndexIdleTimeout, func() { //logger.Logger.Infof("关键词搜索索引空闲超过 %v,自动释放内存", keywordIndexIdleTimeout) s.mu.Lock() s.unload() s.mu.Unlock() })}
// stopIdleTimer 停止空闲计时器func (s *keywordWikiSearcher) stopIdleTimer() { if s.idleTimer != nil { s.idleTimer.Stop() s.idleTimer = nil }}
// loadLocked 不加锁加载(调用方需持有 s.mu 写锁)func (s *keywordWikiSearcher) loadLocked() error { if s.index != nil { return nil }
data, err := os.ReadFile(s.indexPath) if err != nil { return fmt.Errorf("搜索索引文件不存在: %w", err) }
var idx wikiIndex if err := json.Unmarshal(data, &idx); err != nil { return fmt.Errorf("解析搜索索引失败: %w", err) }
s.index = &idx s.resetIdleTimer() logger.Logger.Infof("已加载关键词搜索索引: %d 篇文档, %d 个词条", len(idx.Pages), len(idx.TermIndex)) return nil}
// Search 搜索 Wiki 文档。索引未构建时返回 error。func (s *keywordWikiSearcher) search(query string, maxResults int) ([]wikiResult, error) { s.mu.Lock() if s.index == nil { if err := s.loadLocked(); err != nil { s.mu.Unlock() return nil, fmt.Errorf("搜索索引未构建,请先构建索引: %w", err) } } idx := s.index s.resetIdleTimer() // 每次 Search 重置空闲计时器 s.mu.Unlock()
queryTerms := tokenizeQuery(query) scores := make(map[string]float64) matchedTerms := make(map[string][]string)
queryLower := strings.ToLower(strings.TrimSpace(query))
// TF 评分 for term := range queryTerms { if docs, ok := idx.TermIndex[term]; ok { for filename, tfScore := range docs { scores[filename] += tfScore matchedTerms[filename] = append(matchedTerms[filename], term) } } }
// 标题匹配加权 for filename, info := range idx.Pages { titleLower := strings.ToLower(info.Title) if strings.Contains(titleLower, queryLower) { scores[filename] += 5.0 matchedTerms[filename] = append(matchedTerms[filename], "标题精确匹配") } for term := range queryTerms { if utf8.RuneCountInString(term) >= 3 && strings.Contains(titleLower, term) { scores[filename] += 1.0 } } }
// 排序 type scoredFile struct { filename string score float64 } ranked := make([]scoredFile, 0, len(scores)) for filename, score := range scores { ranked = append(ranked, scoredFile{filename, score}) } sort.Slice(ranked, func(i, j int) bool { return ranked[i].score > ranked[j].score })
// 取 Top N results := make([]wikiResult, 0, maxResults) for _, sf := range ranked { if sf.score <= 0 { continue } info := idx.Pages[sf.filename] if info.ContentLength < 50 { continue } results = append(results, wikiResult{ Title: info.Title, Filename: info.Filename, Categories: info.Categories, Content: info.Content, ContentLen: info.ContentLength, Score: roundFloat(sf.score, 2), }) if len(results) >= maxResults { break } }
return results, nil}
// ========== 分词 ==========
func tokenizeQuery(query string) map[string]struct{} { terms := make(map[string]struct{})
// 中文连续 2-4 字 for _, chunk := range findChineseChunks(query) { runes := []rune(chunk) for i := 0; i < len(runes); i++ { for length := 2; length <= 4; length++ { if i+length <= len(runes) { terms[string(runes[i:i+length])] = struct{}{} } } } }
// 英文词 for _, word := range regexp.MustCompile(`[a-zA-Z0-9]{2,}`).FindAllString(strings.ToLower(query), -1) { terms[word] = struct{}{} }
// 完整查询词 terms[strings.ToLower(strings.TrimSpace(query))] = struct{}{}
return terms}
func findChineseChunks(text string) []string { var chunks []string var current []rune for _, r := range text { if unicode.Is(unicode.Han, r) { current = append(current, r) } else { if len(current) > 0 { chunks = append(chunks, string(current)) current = nil } } } if len(current) > 0 { chunks = append(chunks, string(current)) } return chunks}
// ========== 上下文格式化 ==========
// maxContextTokens 默认最大上下文 token 数const maxContextTokens = 8000
// formatWikiContext 将搜索结果格式化为 AI 聊天的参考上下文func formatWikiContext(results []wikiResult, maxTokens int) string { if len(results) == 0 { return "" }
if maxTokens <= 0 { maxTokens = maxContextTokens }
var parts []string parts = append(parts, "# 饥荒 Wiki 参考文档")
totalChars := 0 maxChars := maxTokens * 2 // 粗略按字符数/2 估算 token
for i, r := range results { header := formatResultHeader(i+1, r) body := r.Content
entryChars := utf8.RuneCountInString(header) + utf8.RuneCountInString(body) + 10 if totalChars+entryChars > maxChars { remaining := maxChars - totalChars - utf8.RuneCountInString(header) - 50 if remaining > 200 { runes := []rune(body) if len(runes) > remaining { body = string(runes[:remaining]) + "\n\n(内容已截断...)" } } else { break } }
parts = append(parts, "") parts = append(parts, header) parts = append(parts, "") parts = append(parts, body)
totalChars += entryChars }
return strings.Join(parts, "\n")}
func formatResultHeader(index int, r wikiResult) string { h := "## " + itoa(index) + ". " + r.Title if len(r.Categories) > 0 { h += " [" + strings.Join(r.Categories, ", ") + "]" } return h}
func itoa(n int) string { if n <= 0 { return "0" } digits := make([]byte, 0, 10) for n > 0 { digits = append([]byte{byte('0' + n%10)}, digits...) n /= 10 } return string(digits)}
func roundFloat(f float64, decimals int) float64 { pow := 1.0 for i := 0; i < decimals; i++ { pow *= 10 } return float64(int(f*pow+0.5)) / pow}
// ================================================================// 构建关键词搜索索引// ================================================================
// BuildIndex 构建关键词搜索索引。扫描所有 .md 文件,构建 TF 评分索引。// force 为 true 时强制重建,否则索引已存在则跳过。func (s *keywordWikiSearcher) buildIndex(force bool) error { s.mu.Lock() defer s.mu.Unlock()
// 非强制模式下,如果索引文件已存在,直接加载 if !force { if _, err := os.Stat(s.indexPath); err == nil { logger.Logger.Infof("关键词索引已存在,加载中...") return s.loadLocked() } }
logger.Logger.Infof("正在构建关键词搜索索引...") start := time.Now()
idx := &wikiIndex{ Pages: make(map[string]wikiPageInfo), CategoryIndex: make(map[string][]string), TermIndex: make(map[string]map[string]float64), }
entries, err := os.ReadDir(s.pagesDir) if err != nil { return fmt.Errorf("扫描 Wiki 页面目录失败: %w", err) }
var mdFiles []string for _, entry := range entries { if !entry.IsDir() && filepath.Ext(entry.Name()) == ".md" { mdFiles = append(mdFiles, filepath.Join(s.pagesDir, entry.Name())) } } sort.Strings(mdFiles) total := len(mdFiles)
for i, fp := range mdFiles { if i%200 == 0 { logger.Logger.Infof("处理中... %d/%d", i, total) }
filename := filepath.Base(fp) content, readErr := os.ReadFile(fp) if readErr != nil { logger.Logger.Warnf("读取 %s 失败: %v", filename, readErr) continue } rawContent := string(content)
// 提取标题 title := extractWikiTitle(rawContent, filename)
// 提取分类 categories := extractWikiCategories(rawContent)
// 提取正文 body := stripWikiMeta(rawContent)
// 提取内部链接 links := extractWikiLinks(body)
// 分词 textForSearch := title + " " + body textForSearch = cleanForTokenize(textForSearch) allTerms := tokenizeText(textForSearch)
// 记录页面信息 idx.Pages[filename] = wikiPageInfo{ Title: title, Categories: categories, Content: body, ContentLength: utf8.RuneCountInString(body), Links: links, Filename: filename, Terms: sortedTermSlice(allTerms), }
// 分类索引 for _, cat := range categories { idx.CategoryIndex[cat] = append(idx.CategoryIndex[cat], filename) }
// 词索引 (TF scoring: 词在文档中出现越多,分越高) textLower := strings.ToLower(textForSearch) termCounts := make(map[string]int) for term := range allTerms { termCounts[term] = strings.Count(textLower, strings.ToLower(term)) }
maxCount := 1 for _, count := range termCounts { if count > maxCount { maxCount = count } }
for term, count := range termCounts { // 归一化 TF + 标题中出现加权 score := float64(count) / float64(maxCount) if strings.Contains(strings.ToLower(title), strings.ToLower(term)) { score *= 2.0 } if idx.TermIndex[term] == nil { idx.TermIndex[term] = make(map[string]float64) } idx.TermIndex[term][filename] = score } }
// 保存索引 if err := saveKeywordIndex(s.indexPath, idx); err != nil { return err }
s.index = idx s.resetIdleTimer()
elapsed := time.Since(start) logger.Logger.Infof("索引构建完成! %d 个页面, %d 个词条, 耗时 %.1fs", total, len(idx.TermIndex), elapsed.Seconds()) return nil}
// ========== 索引构建辅助函数 ==========
// tokenizeText 对文本全文分词(用于构建索引,不含完整查询词)func tokenizeText(text string) map[string]struct{} { terms := make(map[string]struct{})
// 中文连续 2-4 字 for _, chunk := range findChineseChunks(text) { runes := []rune(chunk) for j := 0; j < len(runes); j++ { for length := 2; length <= 4; length++ { if j+length <= len(runes) { terms[string(runes[j:j+length])] = struct{}{} } } } }
// 英文词 for _, word := range regexp.MustCompile(`[a-zA-Z0-9]{2,}`).FindAllString(strings.ToLower(text), -1) { terms[word] = struct{}{} }
return terms}
// cleanForTokenize 去除 markdown 格式字符,压缩空白func cleanForTokenize(text string) string { text = regexp.MustCompile(`[*#>`+"`"+`\[\]()!_~|]`).ReplaceAllString(text, " ") text = regexp.MustCompile(`\s+`).ReplaceAllString(text, " ") return text}
// extractWikiLinks 提取 Wiki 内部链接的显示文本var wikiLinkRe = regexp.MustCompile(`\[([^\]]+)\]\(([^)]+\.md)\)`)
func extractWikiLinks(body string) []string { matches := wikiLinkRe.FindAllStringSubmatch(body, -1) links := make([]string, 0, len(matches)) for _, m := range matches { if len(m) >= 2 { links = append(links, m[1]) } } return links}
// sortedTermSlice 将 term set 转为排序的 slicefunc sortedTermSlice(terms map[string]struct{}) []string { result := make([]string, 0, len(terms)) for t := range terms { result = append(result, t) } sort.Strings(result) return result}
// saveKeywordIndex 保存搜索索引到磁盘func saveKeywordIndex(indexPath string, idx *wikiIndex) error { dir := filepath.Dir(indexPath) if err := os.MkdirAll(dir, 0755); err != nil { return fmt.Errorf("创建索引目录失败: %w", err) }
data, err := json.Marshal(idx) if err != nil { return fmt.Errorf("序列化搜索索引失败: %w", err) }
if err := os.WriteFile(indexPath, data, 0644); err != nil { return fmt.Errorf("保存搜索索引失败: %w", err) }
return nil}文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!
相关文章智能推荐
1
DMP代搭建的小弟跑路了
垃圾堆跑了啊!!!
2
饥荒管理平台ios应用
垃圾堆真的好喜欢Liquid Glass啊啊啊啊啊啊
3
waline独立部署踩坑实录
中转站记录一次waline独立部署
随机文章随机推荐













